Merge pull request #1676 from CREDO23/fix/hitl-parent-interrupt-routing
Some checks are pending
Build and Push Docker Images / build (./surfsense_backend, cu128, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-latest, linux/amd64, amd64, -cuda, production, true, cuda) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_web, cpu, ./surfsense_web/Dockerfile, web, surfsense-web, ubuntu-24.04-arm, linux/arm64, arm64, , runner, false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_web, cpu, ./surfsense_web/Dockerfile, web, surfsense-web, ubuntu-latest, linux/amd64, amd64, , runner, false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / compute_version (push) Waiting to run
Build and Push Docker Images / build (./docker/sandbox, cpu, ./docker/sandbox/Dockerfile, sandbox, surfsense-sandbox, ubuntu-24.04-arm, linux/arm64, arm64, , false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / build (./docker/sandbox, cpu, ./docker/sandbox/Dockerfile, sandbox, surfsense-sandbox, ubuntu-latest, linux/amd64, amd64, , false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_backend, cpu, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-24.04-arm, linux/arm64, arm64, , production, false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_backend, cpu, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-latest, linux/amd64, amd64, , production, false, cpu) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_backend, cu126, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-24.04-arm, linux/arm64, arm64, -cuda126, production, true, cuda126) (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_backend, cu126, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-latest, linux/amd64, amd64, -cuda126, production, true, cuda126) (push) Blocked by required conditions
Build and Push Docker Images / verify_digests (push) Blocked by required conditions
Build and Push Docker Images / create_manifest (backend, surfsense-backend, , cpu) (push) Blocked by required conditions
Build and Push Docker Images / create_manifest (backend, surfsense-backend, -cuda, cuda) (push) Blocked by required conditions
Build and Push Docker Images / create_manifest (backend, surfsense-backend, -cuda126, cuda126) (push) Blocked by required conditions
Build and Push Docker Images / create_manifest (sandbox, surfsense-sandbox, , cpu) (push) Blocked by required conditions
Build and Push Docker Images / create_manifest (web, surfsense-web, , cpu) (push) Blocked by required conditions
Build and Push Docker Images / finalize_release (push) Blocked by required conditions
Build and Push Docker Images / build (./surfsense_backend, cu128, ./surfsense_backend/Dockerfile, backend, surfsense-backend, ubuntu-24.04-arm, linux/arm64, arm64, -cuda, production, true, cuda) (push) Blocked by required conditions

[Fix] Multi-agent HITL reliability + OpenRouter truncation handling
This commit is contained in:
Thierry CH. 2026-08-13 08:07:42 +02:00 committed by GitHub
commit e90ceba275
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
41 changed files with 1850 additions and 142 deletions

View file

@ -34,23 +34,13 @@ def slice_decisions_by_tool_call(
) -> dict[str, dict[str, Any]]:
"""Slice ``decisions`` into ``{tool_call_id: {"decisions": <slice>}}``.
Args:
decisions: Flat list of decisions in the order the SSE stream rendered
them.
pending: Ordered ``(tool_call_id, action_count)`` pairs in the same
order. The slicer consumes ``decisions`` left-to-right.
Returns:
Per-``tool_call_id`` payload dict ready to be written to
``configurable["surfsense_resume_value"]``.
Raises:
ValueError: When the total expected action count differs from the
number of decisions provided. We fail loud rather than silently
dropping or padding so a frontend/backend contract drift surfaces
immediately.
Routes by identity when every decision carries a ``tool_call_id``, else by
position in ``pending`` order. Raises on any count or id mismatch.
"""
pending_list = list(pending)
if decisions and all(d.get("tool_call_id") for d in decisions):
return _route_by_id(decisions, pending_list)
expected = sum(count for _, count in pending_list)
if expected != len(decisions):
raise ValueError(
@ -66,6 +56,34 @@ def slice_decisions_by_tool_call(
return routed
def _route_by_id(
decisions: list[dict[str, Any]],
pending_list: list[tuple[str, int]],
) -> dict[str, dict[str, Any]]:
"""Route id-stamped decisions to their pending tool call, validating identity."""
grouped: dict[str, list[dict[str, Any]]] = {}
for decision in decisions:
grouped.setdefault(str(decision["tool_call_id"]), []).append(decision)
pending_ids = {tool_call_id for tool_call_id, _ in pending_list}
if set(grouped) != pending_ids:
raise ValueError(
"Decision routing mismatch: decisions target "
f"{sorted(grouped)} but pending tool calls are {sorted(pending_ids)}."
)
routed: dict[str, dict[str, Any]] = {}
for tool_call_id, action_count in pending_list:
slice_ = grouped[tool_call_id]
if len(slice_) != action_count:
raise ValueError(
f"Decision count mismatch for tool_call_id={tool_call_id!r}: "
f"expected {action_count} action(s) but received {len(slice_)}."
)
routed[tool_call_id] = {"decisions": slice_}
return routed
def collect_pending_tool_calls(state: Any) -> list[tuple[str, int]]:
"""Extract ``[(tool_call_id, action_count), ...]`` from a paused parent state.
@ -138,6 +156,61 @@ def collect_pending_tool_calls(state: Any) -> list[tuple[str, int]]:
return pending
def collect_pending_parent_interrupts(state: Any) -> list[tuple[str, int]]:
"""Ordered ``(interrupt_id, action_count)`` for unstamped parent-graph interrupts.
Complements :func:`collect_pending_tool_calls`: main-agent
``PermissionMiddleware`` asks and ``DoomLoopMiddleware`` pauses carry no
``tool_call_id`` (they never cross a ``task`` call), so they must be routed
by ``Interrupt.id`` instead. ``action_count`` defaults to 1 for scalar
payloads (e.g. doom-loop) with no ``action_requests``.
"""
pending: list[tuple[str, int]] = []
for interrupt_obj in getattr(state, "interrupts", ()) or ():
value = getattr(interrupt_obj, "value", None)
if not isinstance(value, dict):
continue
if isinstance(value.get("tool_call_id"), str):
continue # subagent-routed; owned by collect_pending_tool_calls
interrupt_id = getattr(interrupt_obj, "id", None)
if not isinstance(interrupt_id, str):
continue
action_requests = value.get("action_requests")
count = (
len(action_requests)
if isinstance(action_requests, list) and action_requests
else 1
)
pending.append((interrupt_id, count))
return pending
def build_parent_resume_map(
decisions: list[dict[str, Any]],
parent_pending: list[tuple[str, int]],
) -> dict[str, Any]:
"""Map ``Interrupt.id → resume_value`` for parent-side interrupts.
Single-action asks deliver the raw decision dict (the site's
``interrupt()`` return); parent sites read it directly and never unwrap a
``{"decisions": [...]}`` bundle. Raises on a decision-count mismatch.
"""
expected = sum(count for _, count in parent_pending)
if expected != len(decisions):
raise ValueError(
f"Decision count mismatch: parent-side interrupts expect "
f"{expected} actions but received {len(decisions)} decisions."
)
out: dict[str, Any] = {}
cursor = 0
for interrupt_id, count in parent_pending:
chunk = decisions[cursor : cursor + count]
cursor += count
out[interrupt_id] = chunk[0] if count == 1 else {"decisions": chunk}
return out
def build_lg_resume_map(
state: Any, by_tool_call_id: dict[str, dict[str, Any]]
) -> dict[str, dict[str, Any]]:

View file

@ -13,15 +13,15 @@ import json
import logging
import time
from collections.abc import Awaitable, Callable
from typing import Annotated, Any, NoReturn, TypeVar
from typing import Annotated, Any, Literal, NoReturn, TypeVar
from deepagents.middleware.subagents import TASK_TOOL_DESCRIPTION
from langchain.tools import BaseTool, ToolRuntime
from langchain_core.messages import HumanMessage, ToolMessage
from langchain_core.runnables import Runnable
from langchain_core.tools import StructuredTool
from langgraph.errors import GraphInterrupt
from langgraph.types import Command, Interrupt
from pydantic import ConfigDict, Field, create_model, field_validator
from app.agents.chat.multi_agent_chat.constants import LEGACY_SUBAGENT_ALIASES
from app.agents.chat.multi_agent_chat.subagents.shared.invocation import (
@ -58,6 +58,24 @@ from .spawn_paused import is_spawn_paused
logger = logging.getLogger(__name__)
_perf_log = get_perf_logger()
_DEFAULT_TASK_DESCRIPTION = (
"Invoke a specialist subagent: pass its `subagent_type` and a full "
"`description` of the task. See the `<specialists>` roster for who exists."
)
_MISSING_RUNTIME_ERROR = (
"task: could not read the tool runtime for this call (likely a truncated or "
"malformed tool call). Re-issue the task with a complete `description` and "
"`subagent_type`."
)
def _runtime_error(runtime: ToolRuntime | None) -> str | None:
"""Model-readable error string if no usable runtime was injected, else ``None``."""
if runtime is None or not getattr(runtime, "tool_call_id", None):
return _MISSING_RUNTIME_ERROR
return None
class SubagentInvokeTimeoutError(Exception):
"""Raised when ``subagent.ainvoke`` exceeds the configured wall-clock budget.
@ -194,18 +212,7 @@ def build_task_tool_with_parent_config(
for spec in subagents
if (provider := spec.get(SURF_CONTEXT_HINT_PROVIDER_KEY)) is not None
}
subagent_description_str = "\n".join(
f"- {s['name']}: {s['description']}" for s in subagents
)
if task_description is None:
description = TASK_TOOL_DESCRIPTION.format(
available_agents=subagent_description_str
)
elif "{available_agents}" in task_description:
description = task_description.format(available_agents=subagent_description_str)
else:
description = task_description
description = task_description or _DEFAULT_TASK_DESCRIPTION
def _billable_call_update(
subagent_type: str, runtime: ToolRuntime
@ -669,6 +676,8 @@ def build_task_tool_with_parent_config(
),
] = None,
) -> str | Command:
if (err := _runtime_error(runtime)) is not None:
return err
if tasks is not None:
return (
"task: batch mode (`tasks=[...]`) is only supported on the async "
@ -687,8 +696,6 @@ def build_task_tool_with_parent_config(
f"We cannot invoke subagent {subagent_type} because it does not exist, "
f"the only allowed types are {allowed_types}"
)
if not runtime.tool_call_id:
raise ValueError("Tool call ID is required for subagent invocation")
subagent, subagent_state = _validate_and_prepare_state(
subagent_type, description, runtime
)
@ -850,6 +857,8 @@ def build_task_tool_with_parent_config(
] = None,
) -> str | Command:
atask_start = time.perf_counter()
if (err := _runtime_error(runtime)) is not None:
return err
# Ops kill switch: short-circuit every task() call for this workspace
# so the orchestrator stops hammering downstream APIs.
if await is_spawn_paused(workspace_id):
@ -869,8 +878,6 @@ def build_task_tool_with_parent_config(
"task: cannot combine `tasks` with `description`/`subagent_type`. "
"Use either single-mode (description+subagent_type) or batch-mode (tasks)."
)
if not runtime.tool_call_id:
raise ValueError("Tool call ID is required for subagent invocation")
coerced = _coerce_batch_arg(tasks)
if isinstance(coerced, str):
return coerced
@ -897,8 +904,6 @@ def build_task_tool_with_parent_config(
f"We cannot invoke subagent {subagent_type} because it does not exist, "
f"the only allowed types are {allowed_types}"
)
if not runtime.tool_call_id:
raise ValueError("Tool call ID is required for subagent invocation")
subagent, subagent_state = _validate_and_prepare_state(
subagent_type, description, runtime
)
@ -1130,4 +1135,72 @@ def build_task_tool_with_parent_config(
func=task,
coroutine=atask,
description=description,
args_schema=_build_task_args_schema(subagent_names),
handle_validation_error=_on_invalid_task_args(subagent_names),
)
def _build_task_args_schema(subagent_names: set[str]) -> type:
"""Args schema constraining single-mode ``subagent_type`` to the live roster.
The ``Literal`` surfaces the roster as a provider-side enum. A before-validator
rewrites legacy connector aliases onto their consolidated route so paused
pre-consolidation checkpoints still resolve; the aliases stay out of the enum.
"""
roster = sorted(subagent_names)
def _canonicalize_legacy(cls, value):
if isinstance(value, str) and value not in subagent_names:
return LEGACY_SUBAGENT_ALIASES.get(value, value)
return value
return create_model(
"TaskToolArgs",
__config__=ConfigDict(arbitrary_types_allowed=True),
__validators__={
"_canonicalize_legacy_subagent_type": field_validator(
"subagent_type", mode="before"
)(_canonicalize_legacy)
},
# ``runtime`` is injected by ToolNode; it must be a declared field or
# validation drops it (ToolRuntime is directly-injected, so it is not in
# ``_injected_args_keys``). Bare ``ToolRuntime`` (not a Union) so
# ``_is_directly_injected_arg_type`` keeps it out of the model-facing schema.
runtime=(ToolRuntime, Field(default=None)),
description=(
str | None,
Field(
default=None,
description="Single-mode: full task prompt. Required unless `tasks` is provided.",
),
),
subagent_type=(
Literal[tuple(roster)] | None if roster else str | None,
Field(
default=None,
description="Single-mode: which specialist to invoke. Required unless `tasks` is provided.",
),
),
tasks=(
list[dict] | None,
Field(
default=None,
description=(
"Batch-mode: array of `{description, subagent_type}` objects to "
"fan out concurrently. Mutually exclusive with single-mode args."
),
),
),
)
def _on_invalid_task_args(subagent_names: set[str]) -> Callable[[Exception], str]:
allowed = ", ".join(f"`{n}`" for n in sorted(subagent_names))
def _handle(_exc: Exception) -> str:
return (
f"Invalid `task` arguments. `subagent_type` must be one of: {allowed}. "
"For batch mode send `tasks=[{description, subagent_type}, ...]`."
)
return _handle

View file

@ -0,0 +1,190 @@
"""Auto-continue a final text answer cut off by ``max_tokens``.
``langchain_litellm`` drops ``finish_reason`` from streamed chunks, so a
token-limit cut would otherwise reach the user as a silent stub. When the last
message is a tool-call-free text answer that hit its output cap, re-invoke with
the partial prefilled and stitch the pieces until it finishes. Tool-call
truncation (invalid partial JSON) and non-string content fall back to the
usage-based truncation marker.
"""
from __future__ import annotations
import dataclasses
from typing import TYPE_CHECKING, Any
from langchain.agents.middleware import AgentMiddleware
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from app.services.token_tracking_service import (
get_current_accumulator,
is_output_truncated,
)
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
_CONTINUE_NUDGE = HumanMessage(
content=(
"Your previous response was cut off before it finished. Continue exactly "
"where you left off. Do not repeat any text you already wrote and do not "
"add any preamble."
)
)
def _final_text_ai(response: Any) -> AIMessage | None:
result = getattr(response, "result", None) or []
msg = result[-1] if result else None
return msg if isinstance(msg, AIMessage) else None
def _plain_text(ai: AIMessage) -> str | None:
return ai.content if isinstance(ai.content, str) else None
def _output_tokens(ai: AIMessage) -> int:
return (ai.usage_metadata or {}).get("output_tokens", 0) or 0
def _finish_reason(ai: AIMessage) -> str | None:
return (ai.response_metadata or {}).get("finish_reason")
def _model_max_tokens(request: Any) -> int | None:
return getattr(getattr(request, "model", None), "max_tokens", None)
def _continuation_messages(base: list[BaseMessage], accumulated: str) -> list[BaseMessage]:
return [*base, AIMessage(content=accumulated), _CONTINUE_NUDGE]
def _merge(first_ai: AIMessage, text: str, total_out: int) -> AIMessage:
usage = dict(first_ai.usage_metadata or {})
if usage:
usage["output_tokens"] = total_out
usage["total_tokens"] = usage.get("input_tokens", 0) + total_out
metadata = dict(first_ai.response_metadata or {})
metadata["finish_reason"] = "stop"
return AIMessage(
content=text,
id=first_ai.id,
usage_metadata=usage or None,
response_metadata=metadata,
additional_kwargs=first_ai.additional_kwargs,
)
class ContinueOnMaxLengthMiddleware(AgentMiddleware): # type: ignore[type-arg]
"""Stitch continuations onto a truncated, tool-call-free text answer."""
def __init__(self, max_continuations: int = 2) -> None:
super().__init__()
self.max_continuations = max_continuations
def _should_continue(self, ai: AIMessage, max_tokens: int | None) -> bool:
return not ai.tool_calls and is_output_truncated(
_finish_reason(ai), _output_tokens(ai), max_tokens
)
def _finalize(
self,
*,
response: Any,
last_response: Any,
first_ai: AIMessage,
ai: AIMessage | None,
accumulated: str,
total_out: int,
done: int,
max_tokens: int | None,
) -> Any:
if done == 0:
return response
# Recovered a complete answer: the per-call ``length`` the token callback
# flagged is no longer user-visible, so clear the marker.
if ai is not None and not self._should_continue(ai, max_tokens):
acc = get_current_accumulator()
if acc is not None:
acc.truncated = False
return dataclasses.replace(
last_response, result=[_merge(first_ai, accumulated, total_out)]
)
def wrap_model_call( # type: ignore[override]
self,
request: Any,
handler: Callable[[Any], Any],
) -> Any:
response = handler(request)
ai = _final_text_ai(response)
max_tokens = _model_max_tokens(request)
if ai is None or (accumulated := _plain_text(ai)) is None:
return response
total_out = _output_tokens(ai)
first_ai, last_response, done = ai, response, 0
while done < self.max_continuations and self._should_continue(ai, max_tokens):
done += 1
last_response = handler(
request.override(messages=_continuation_messages(request.messages, accumulated))
)
ai = _final_text_ai(last_response)
if ai is None or (piece := _plain_text(ai)) is None:
break
accumulated += piece
total_out += _output_tokens(ai)
return self._finalize(
response=response,
last_response=last_response,
first_ai=first_ai,
ai=ai,
accumulated=accumulated,
total_out=total_out,
done=done,
max_tokens=max_tokens,
)
async def awrap_model_call( # type: ignore[override]
self,
request: Any,
handler: Callable[[Any], Awaitable[Any]],
) -> Any:
response = await handler(request)
ai = _final_text_ai(response)
max_tokens = _model_max_tokens(request)
if ai is None or (accumulated := _plain_text(ai)) is None:
return response
total_out = _output_tokens(ai)
first_ai, last_response, done = ai, response, 0
while done < self.max_continuations and self._should_continue(ai, max_tokens):
done += 1
last_response = await handler(
request.override(messages=_continuation_messages(request.messages, accumulated))
)
ai = _final_text_ai(last_response)
if ai is None or (piece := _plain_text(ai)) is None:
break
accumulated += piece
total_out += _output_tokens(ai)
return self._finalize(
response=response,
last_response=last_response,
first_ai=first_ai,
ai=ai,
accumulated=accumulated,
total_out=total_out,
done=done,
max_tokens=max_tokens,
)
def build_continue_on_max_length_mw(flags: Any) -> ContinueOnMaxLengthMiddleware | None:
from ...shared.middleware.flags import enabled
if not enabled(flags, "enable_continue_on_max_length"):
return None
return ContinueOnMaxLengthMiddleware(max_continuations=2)

View file

@ -12,8 +12,9 @@ the agent has likely entered an infinite loop. We surface this to the
user as an interrupt with ``permission="doom_loop"`` so the UI can
render an "Are you stuck? Continue / cancel?" affordance.
This ships **OFF by default** until the frontend explicitly handles
``context.permission == "doom_loop"`` interrupts.
Ships ON by default (``enable_doom_loop``): the frontend renders a dedicated
continue/stop card for ``context.permission == "doom_loop"`` and the resume
router routes the reply back by ``Interrupt.id``.
Wire format: uses SurfSense's existing ``interrupt()`` payload shape
(see ``app/agents/shared/tools/hitl.py``):

View file

@ -85,6 +85,7 @@ from .checkpointed_subagent_middleware.task_description import (
TASK_TOOL_DESCRIPTION,
)
from .context_editing import build_context_editing_mw
from .continue_on_max_length import build_continue_on_max_length_mw
from .dedup_hitl import build_dedup_hitl_mw
from .doom_loop import build_doom_loop_mw
from .kb_persistence import build_kb_persistence_mw
@ -273,6 +274,7 @@ def build_main_agent_deepagent_middleware(
task_description=TASK_TOOL_DESCRIPTION,
workspace_id=workspace_id,
),
build_continue_on_max_length_mw(flags),
resilience.model_call_limit,
resilience.tool_call_limit,
build_context_editing_mw(

View file

@ -11,28 +11,7 @@ from app.agents.chat.multi_agent_chat.shared.middleware.flags import enabled
from .middleware import ToolCallNameRepairMiddleware
# deepagents-built-in tool names the repair pass treats as known.
_DEEPAGENT_BUILTIN_TOOL_NAMES: frozenset[str] = frozenset(
{
"write_todos",
"ls",
"read_file",
"write_file",
"edit_file",
"glob",
"grep",
"execute",
"task",
"mkdir",
"cd",
"pwd",
"move_file",
"rm",
"rmdir",
"list_tree",
"execute_code",
}
)
_MIDDLEWARE_BOUND_TOOL_NAMES: frozenset[str] = frozenset({"task", "write_todos"})
def build_repair_mw(
@ -43,7 +22,7 @@ def build_repair_mw(
if not enabled(flags, "enable_tool_call_repair"):
return None
registered_names: set[str] = {t.name for t in tools}
registered_names |= _DEEPAGENT_BUILTIN_TOOL_NAMES
registered_names |= _MIDDLEWARE_BOUND_TOOL_NAMES
return ToolCallNameRepairMiddleware(
registered_tool_names=registered_names,
fuzzy_match_threshold=None,

View file

@ -44,6 +44,8 @@ class AgentFeatureFlags:
enable_tool_call_limit: bool = True
enable_tool_call_repair: bool = True
enable_doom_loop: bool = True
# Stitch continuations onto a final answer cut off by max_tokens.
enable_continue_on_max_length: bool = True
# Safety — permissions, concurrency, tool-set narrowing
enable_permission: bool = True
@ -105,6 +107,7 @@ class AgentFeatureFlags:
enable_tool_call_limit=False,
enable_tool_call_repair=False,
enable_doom_loop=False,
enable_continue_on_max_length=False,
enable_permission=False,
enable_busy_mutex=False,
enable_llm_tool_selector=False,
@ -134,6 +137,9 @@ class AgentFeatureFlags:
"SURFSENSE_ENABLE_TOOL_CALL_REPAIR", True
),
enable_doom_loop=_env_bool("SURFSENSE_ENABLE_DOOM_LOOP", True),
enable_continue_on_max_length=_env_bool(
"SURFSENSE_ENABLE_CONTINUE_ON_MAX_LENGTH", True
),
# Safety
enable_permission=_env_bool("SURFSENSE_ENABLE_PERMISSION", True),
enable_busy_mutex=_env_bool("SURFSENSE_ENABLE_BUSY_MUTEX", True),
@ -176,6 +182,7 @@ class AgentFeatureFlags:
self.enable_tool_call_limit,
self.enable_tool_call_repair,
self.enable_doom_loop,
self.enable_continue_on_max_length,
self.enable_permission,
self.enable_busy_mutex,
self.enable_llm_tool_selector,

View file

@ -36,4 +36,46 @@ def build_deny_message(tool_call: dict[str, Any], rule: Rule) -> ToolMessage:
)
__all__ = ["build_deny_message"]
def build_reject_message(tool_call: dict[str, Any]) -> ToolMessage:
"""Reject without feedback: model must stop retrying this call and ask the user."""
err = StreamingError(
code="permission_denied",
retryable=False,
suggestion="Do not retry this call; ask the user how to proceed.",
)
return ToolMessage(
content=(
f"The user rejected tool {tool_call.get('name')!r}. Do not retry the "
"same call; ask the user how they would like to proceed."
),
tool_call_id=tool_call.get("id") or "",
name=tool_call.get("name"),
status="error",
additional_kwargs={"error": err.model_dump()},
)
def build_correction_message(tool_call: dict[str, Any], feedback: str) -> ToolMessage:
"""Reject with feedback: surface the correction so the model can retry differently."""
err = StreamingError(
code="permission_denied",
retryable=True,
suggestion="Adjust the call per the user's feedback and try again.",
)
return ToolMessage(
content=(
f"The user rejected tool {tool_call.get('name')!r} with feedback: "
f"{feedback}"
),
tool_call_id=tool_call.get("id") or "",
name=tool_call.get("name"),
status="error",
additional_kwargs={"error": err.model_dump()},
)
__all__ = [
"build_correction_message",
"build_deny_message",
"build_reject_message",
]

View file

@ -27,12 +27,11 @@ from langchain_core.tools import BaseTool
from langgraph.runtime import Runtime
from app.agents.chat.multi_agent_chat.shared.permissions.model import Ruleset
from app.agents.chat.runtime.errors import CorrectedError, RejectedError
from app.services.user_tool_allowlist import TrustedToolSaver
from ..ask.edit import merge_edited_args
from ..ask.request import request_permission_decision
from ..deny import build_deny_message
from ..deny import build_correction_message, build_deny_message, build_reject_message
from .evaluation import evaluate_tool_call
from .pattern_resolver import PatternResolver
from .ruleset_view import all_rulesets
@ -173,15 +172,16 @@ class PermissionMiddleware(AgentMiddleware): # type: ignore[type-arg]
elif kind == "reject":
feedback = decision.get("feedback")
if isinstance(feedback, str) and feedback.strip():
raise CorrectedError(feedback, tool=name)
raise RejectedError(
tool=name, pattern=patterns[0] if patterns else None
)
deny_messages.append(build_correction_message(call, feedback))
else:
deny_messages.append(build_reject_message(call))
any_change = True
else:
logger.warning(
"Unknown permission decision %r; treating as reject", kind
)
raise RejectedError(tool=name)
deny_messages.append(build_reject_message(call))
any_change = True
continue
kept_calls.append(call)

View file

@ -20,6 +20,11 @@ from .prompts import load_readonly_description
TOOL_NAME = "ask_knowledge_base"
_MISSING_TOOL_CALL_ID_ERROR = (
"Error: ask_knowledge_base was invoked without a tool call id and cannot "
"run. Retry the call as a normal tool call."
)
def _forward_state(runtime: ToolRuntime, query: str) -> dict:
forwarded = {k: v for k, v in runtime.state.items() if k not in EXCLUDED_STATE_KEYS}
@ -84,7 +89,7 @@ def build_ask_knowledge_base_tool(
runtime: ToolRuntime,
) -> str | Command:
if not runtime.tool_call_id:
raise ValueError("Tool call ID is required for ask_knowledge_base")
return _MISSING_TOOL_CALL_ID_ERROR
sub_state = _forward_state(runtime, query)
sub_config = subagent_invoke_config(runtime)
result = _resolve().invoke(sub_state, config=sub_config)
@ -99,7 +104,7 @@ def build_ask_knowledge_base_tool(
runtime: ToolRuntime,
) -> str | Command:
if not runtime.tool_call_id:
raise ValueError("Tool call ID is required for ask_knowledge_base")
return _MISSING_TOOL_CALL_ID_ERROR
sub_state = _forward_state(runtime, query)
sub_config = subagent_invoke_config(runtime)
result = await _resolve().ainvoke(sub_state, config=sub_config)

View file

@ -51,33 +51,6 @@ class StreamingError(BaseModel):
frozen = True
class RejectedError(Exception):
"""Raised when the user rejects a permission ask without feedback.
Caught by :class:`PermissionMiddleware`; the agent stops the current
tool fan-out and surfaces a user-facing rejection.
"""
def __init__(self, *, tool: str | None = None, pattern: str | None = None) -> None:
super().__init__(f"Permission rejected for tool {tool!r}, pattern {pattern!r}")
self.tool = tool
self.pattern = pattern
class CorrectedError(Exception):
"""Raised when the user rejects a permission ask *with* feedback.
The :class:`PermissionMiddleware` translates the feedback into a
synthetic ``ToolMessage`` so the model sees the user's correction
and can retry the request differently.
"""
def __init__(self, feedback: str, *, tool: str | None = None) -> None:
super().__init__(feedback)
self.feedback = feedback
self.tool = tool
class BusyError(Exception):
"""Raised when a second prompt arrives while the same thread is mid-stream."""
@ -88,8 +61,6 @@ class BusyError(Exception):
__all__ = [
"BusyError",
"CorrectedError",
"ErrorCode",
"RejectedError",
"StreamingError",
]

View file

@ -63,6 +63,7 @@ from app.schemas.new_chat import (
NewChatThreadUpdate,
NewChatThreadVisibilityUpdate,
NewChatThreadWithMessages,
PendingInterruptsResponse,
PublicChatSnapshotCreateResponse,
PublicChatSnapshotListResponse,
RegenerateRequest,
@ -223,6 +224,34 @@ def _raise_if_thread_busy_for_start(thread_id: int) -> None:
)
async def _raise_if_thread_awaiting_approval(thread_id: int, checkpointer) -> None:
"""Refuse a fresh turn when the thread's checkpoint has a pending interrupt.
The busy mutex releases on an ``interrupt()`` pause, so a paused thread
reads as idle to ``_raise_if_thread_busy_for_start``. Running ``new_chat`` /
``regenerate`` over that checkpoint would orphan the pending approval the
user must resume or cancel it first. ``resume`` is exempt; it's the path
that clears the pause.
"""
from app.tasks.chat.streaming.helpers.interrupt_inspector import (
pending_interrupt_entries_from_writes,
)
checkpoint_tuple = await checkpointer.aget_tuple(
{"configurable": {"thread_id": str(thread_id)}}
)
if checkpoint_tuple is None:
return
if pending_interrupt_entries_from_writes(checkpoint_tuple.pending_writes):
raise HTTPException(
status_code=409,
detail={
"errorCode": "THREAD_AWAITING_APPROVAL",
"message": chat_error_message("THREAD_AWAITING_APPROVAL"),
},
)
def _find_pre_turn_checkpoint_id(
checkpoint_tuples: list,
*,
@ -1743,6 +1772,11 @@ async def handle_new_chat(
# Check thread-level access based on visibility
await check_thread_access(session, thread, user)
_raise_if_thread_busy_for_start(request.chat_id)
from app.agents.chat.runtime.checkpointer import get_checkpointer
await _raise_if_thread_awaiting_approval(
request.chat_id, await get_checkpointer()
)
filesystem_selection = _resolve_filesystem_selection(
mode=request.filesystem_mode,
client_platform=request.client_platform,
@ -1980,15 +2014,16 @@ async def regenerate_response(
# Check thread-level access based on visibility
await check_thread_access(session, thread, user)
_raise_if_thread_busy_for_start(thread_id)
# Get the checkpointer and state history
checkpointer = await get_checkpointer()
await _raise_if_thread_awaiting_approval(thread_id, checkpointer)
filesystem_selection = _resolve_filesystem_selection(
mode=request.filesystem_mode,
client_platform=request.client_platform,
local_mounts=request.local_filesystem_mounts,
)
# Get the checkpointer and state history
checkpointer = await get_checkpointer()
config = {"configurable": {"thread_id": str(thread_id)}}
# Collect checkpoint tuples from the async iterator
@ -2376,6 +2411,86 @@ async def regenerate_response(
# =============================================================================
@router.get(
"/threads/{thread_id}/pending-interrupts",
response_model=PendingInterruptsResponse,
)
async def get_pending_interrupts(
thread_id: int,
session: AsyncSession = Depends(get_async_session),
auth: AuthContext = Depends(get_auth_context),
):
"""Return the thread's paused HITL interrupts, if any.
The live approval card lives only in the frontend's in-memory stream
overlay, so a page refresh loses it while the LangGraph checkpoint stays
paused. The frontend calls this on thread load to re-render the card and
let the user resume or reject.
"""
from app.agents.chat.runtime.checkpointer import get_checkpointer
from app.services.new_streaming_service import VercelStreamingService
from app.tasks.chat.streaming.helpers.interrupt_inspector import (
pending_interrupt_entries_from_writes,
)
result = await session.execute(
select(NewChatThread).filter(NewChatThread.id == thread_id)
)
thread = result.scalars().first()
if not thread:
raise HTTPException(status_code=404, detail="Thread not found")
await check_permission(
session,
auth,
thread.workspace_id,
Permission.CHATS_READ.value,
"You don't have permission to read chats in this workspace",
)
await check_thread_access(session, thread, user=auth.user)
checkpointer = await get_checkpointer()
checkpoint_tuple = await checkpointer.aget_tuple(
{"configurable": {"thread_id": str(thread_id)}}
)
if checkpoint_tuple is None:
return PendingInterruptsResponse()
entries = pending_interrupt_entries_from_writes(checkpoint_tuple.pending_writes)
if not entries:
return PendingInterruptsResponse()
service = VercelStreamingService()
payloads: list[dict] = []
for value, interrupt_id in entries:
payload = service._normalize_interrupt_payload(value)
if interrupt_id is not None:
payload = {**payload, "interrupt_id": interrupt_id}
payloads.append(payload)
# Reattach the card to the paused turn's assistant row. ``turn_id`` on the
# checkpoint mirrors ``NewChatMessage.turn_id``; fall back to the newest
# assistant row (the paused turn is always the head).
metadata = checkpoint_tuple.metadata or {}
turn_id = metadata.get("turn_id") if isinstance(metadata, dict) else None
assistant_query = (
select(NewChatMessage.id)
.filter(
NewChatMessage.thread_id == thread_id,
NewChatMessage.role == NewChatMessageRole.ASSISTANT,
)
.order_by(NewChatMessage.created_at.desc())
)
if turn_id:
assistant_query = assistant_query.filter(NewChatMessage.turn_id == turn_id)
assistant_message_id = (await session.execute(assistant_query.limit(1))).scalar()
return PendingInterruptsResponse(
assistant_message_id=assistant_message_id,
pending_interrupts=payloads,
)
@router.post("/threads/{thread_id}/resume")
async def resume_chat(
thread_id: int,

View file

@ -425,6 +425,7 @@ class AgentToolInfo(BaseModel):
class ResumeDecision(BaseModel):
type: Literal["approve", "edit", "reject", "approve_always"]
edited_action: dict[str, Any] | None = None
tool_call_id: str | None = None
class ResumeRequest(BaseModel):
@ -458,6 +459,20 @@ class ResumeRequest(BaseModel):
)
class PendingInterruptsResponse(BaseModel):
"""Paused HITL interrupts for a thread, reconstructed from the checkpoint.
Lets the frontend re-render approval cards after a page refresh (the live
``chatStreamStore`` overlay lives only in module RAM). Each payload matches
the ``data-interrupt-request`` SSE ``data`` shape (carries ``interrupt_id``
/ ``tool_call_id``); ``assistant_message_id`` is the paused turn's row so
the card reattaches to the right message.
"""
assistant_message_id: int | None = None
pending_interrupts: list[dict[str, Any]] = Field(default_factory=list)
class CancelActiveTurnResponse(BaseModel):
"""Response for canceling an active turn on a chat thread."""

View file

@ -506,18 +506,25 @@ class VercelStreamingService:
},
)
def format_interrupt_request(self, interrupt_value: dict[str, Any]) -> str:
def format_interrupt_request(
self, interrupt_value: dict[str, Any], *, interrupt_id: str | None = None
) -> str:
"""Format an interrupt request for human-in-the-loop approval.
Args:
interrupt_value: The interrupt payload from either:
- interrupt_on config: {action_requests: [...], review_configs: [...]}
- interrupt() primitive: {type: "...", message: "...", action: {...}, context: {...}}
interrupt_id: langgraph ``Interrupt.id``. The only stable handle for
parent-side interrupts (doom-loop, permission asks) that carry no
``tool_call_id``; the frontend uses it to render and resume them.
Returns:
str: SSE formatted interrupt request data part
"""
normalized_payload = self._normalize_interrupt_payload(interrupt_value)
if interrupt_id is not None:
normalized_payload = {**normalized_payload, "interrupt_id": interrupt_id}
return self.format_data("interrupt-request", normalized_payload)
def _normalize_interrupt_payload(

View file

@ -32,6 +32,25 @@ from app.db import TokenUsage
logger = logging.getLogger(__name__)
def is_output_truncated(
finish_reason: str | None,
completion_tokens: int,
max_tokens: int | None,
) -> bool:
"""True when a generation was cut off by the model's output-token cap.
``finish_reason == "length"`` is authoritative. When it is absent
``langchain_litellm`` drops it from streamed chunks fall back to usage:
hitting the configured ``max_tokens`` is the same event in practice. Any
other explicit reason (``stop``/``tool_calls``/) is a clean finish.
"""
if finish_reason == "length":
return True
if finish_reason:
return False
return bool(max_tokens) and completion_tokens >= max_tokens
def _bare_model_name(model: str) -> str:
"""Return a model identifier with any provider routing prefix stripped.
@ -75,6 +94,8 @@ class TurnTokenAccumulator:
model_metadata_by_bare: dict[str, dict[str, str | None]] = field(
default_factory=dict
)
# Set when any chat call in the turn was cut off by its output-token cap.
truncated: bool = False
def register_model_metadata(
self,
@ -468,6 +489,17 @@ class TokenTrackingCallback(CustomLogger):
call_kind=call_kind,
)
# Streaming drops finish_reason, but the reconstructed response_obj keeps it.
if not is_image:
choices = getattr(response_obj, "choices", None) or []
finish_reason = (
getattr(choices[0], "finish_reason", None) if choices else None
)
if is_output_truncated(
finish_reason, completion_tokens, kwargs.get("max_tokens")
):
acc.truncated = True
# Per-LLM-call wall-clock latency (LiteLLM passes datetime objects).
call_latency_s: float | None = None
try:

View file

@ -30,7 +30,7 @@ from app.tasks.chat.streaming.contract.file_contract import (
)
from app.tasks.chat.streaming.graph_stream.event_stream import stream_output
from app.tasks.chat.streaming.helpers.interrupt_inspector import (
all_interrupt_values,
all_interrupt_entries,
)
from app.tasks.chat.streaming.shared.stream_result import StreamResult
from app.tasks.chat.streaming.shared.utils import safe_float
@ -125,7 +125,8 @@ async def stream_agent_events(
# A turn paused for approval is not a finished turn: the graph resumes into
# this same working copy, so the copy has to outlive the stream.
pending_values = all_interrupt_values(state)
pending_entries = all_interrupt_entries(state)
pending_values = [value for value, _ in pending_entries]
# Same safety net for the git-native path. The pending state is the turn's
# working copy on disk, so no state markers gate it: no copy (or aafter_agent
@ -221,5 +222,7 @@ async def stream_agent_events(
# the resume slicer in
# ``checkpointed_subagent_middleware.resume_routing`` consumes in the
# same order — keeping emit and resume in lock-step.
for interrupt_value in pending_values:
yield streaming_service.format_interrupt_request(interrupt_value)
for interrupt_value, interrupt_id in pending_entries:
yield streaming_service.format_interrupt_request(
interrupt_value, interrupt_id=interrupt_id
)

View file

@ -41,6 +41,10 @@ CHAT_ERROR_MESSAGES: dict[str, str] = {
"seconds or switch models."
),
"SERVER_ERROR": ("We couldn't complete this response right now. Please try again."),
"THREAD_AWAITING_APPROVAL": (
"This thread is waiting on your approval. Respond to the pending action, "
"or stop the response, before sending a new message."
),
"THREAD_BUSY": (
"Another response is still finishing for this thread. Please try again "
"in a moment."

View file

@ -43,18 +43,40 @@ async def build_resume_routing(
"""
from app.agents.chat.multi_agent_chat.main_agent.middleware.checkpointed_subagent_middleware.resume_routing import (
build_lg_resume_map,
build_parent_resume_map,
collect_pending_parent_interrupts,
collect_pending_tool_calls,
slice_decisions_by_tool_call,
)
parent_state = await agent.aget_state({"configurable": {"thread_id": str(chat_id)}})
pending = collect_pending_tool_calls(parent_state)
parent_pending = collect_pending_parent_interrupts(parent_state)
_perf_log.info(
"[hitl_route] resume_entry chat_id=%s decisions=%d pending_subagents=%d",
"[hitl_route] resume_entry chat_id=%s decisions=%d pending_subagents=%d "
"pending_parent=%d",
chat_id,
len(decisions),
len(pending),
len(parent_pending),
)
if parent_pending:
# Parent-side interrupts route by Interrupt.id with no subagent bridge.
# A mix with subagent pauses can't occur (they fire pre-delegation);
# fail loud rather than mis-route.
if pending:
raise ValueError(
"Cannot resume: both parent-side and subagent-side interrupts "
f"are pending (parent={len(parent_pending)}, "
f"subagent={len(pending)}); mixed HITL routing is unsupported."
)
lg_resume_map = build_parent_resume_map(decisions, parent_pending)
return ResumeRoutingPayload(
routed_resume_value={},
lg_resume_map=lg_resume_map,
)
routed_resume_value = slice_decisions_by_tool_call(decisions, pending)
lg_resume_map = build_lg_resume_map(parent_state, routed_resume_value)
return ResumeRoutingPayload(

View file

@ -50,5 +50,6 @@ def iter_token_usage_frame(
"total_tokens": accumulator.grand_total,
"cost_micros": accumulator.total_cost_micros,
"call_details": accumulator.serialized_calls(),
"truncated": accumulator.truncated,
},
)

View file

@ -9,33 +9,45 @@ correlates each frame back to the right paused subagent via the stamped
from __future__ import annotations
from collections.abc import Iterable
from typing import Any
# LangGraph persists interrupts as writes to this channel, one per paused task.
# The named constant is private as of V1.0, but the channel string is the
# durable on-disk value.
_INTERRUPT_CHANNEL = "__interrupt__"
def all_interrupt_values(state: Any) -> list[dict[str, Any]]:
"""Return every interrupt payload across the snapshot, in traversal order.
def all_interrupt_entries(state: Any) -> list[tuple[dict[str, Any], str | None]]:
"""Return ``(value, interrupt_id)`` for every pending interrupt, in order.
Walks ``state.tasks[*].interrupts`` first (langgraph's per-task buckets,
which carry one interrupt per paused subagent) and falls back to
``state.interrupts`` when the per-task lists are empty. Order matches the
snapshot's iteration order so the emit-time order on the SSE stream agrees
with ``collect_pending_tool_calls`` consumption order on resume.
with the resume slicer's consumption order.
The ``interrupt_id`` (langgraph ``Interrupt.id``) is the only stable handle
for parent-side interrupts (doom-loop, permission asks) that carry no
``tool_call_id``; it lets the frontend render and resume them.
Defensive against malformed snapshots: tasks/interrupts that raise on
attribute access are skipped silently. Non-dict values are skipped the
chat-stream contract requires structured interrupt payloads.
"""
def _extract(candidate: Any) -> dict[str, Any] | None:
def _extract(candidate: Any) -> tuple[dict[str, Any], str | None] | None:
if isinstance(candidate, dict):
value = candidate.get("value", candidate)
return value if isinstance(value, dict) else None
value = getattr(candidate, "value", None)
if isinstance(value, dict):
return value
return None
interrupt_id = candidate.get("id")
else:
value = getattr(candidate, "value", None)
interrupt_id = getattr(candidate, "id", None)
if not isinstance(value, dict):
return None
return value, (str(interrupt_id) if interrupt_id is not None else None)
values: list[dict[str, Any]] = []
entries: list[tuple[dict[str, Any], str | None]] = []
saw_task_interrupt = False
for task in getattr(state, "tasks", ()) or ():
@ -48,10 +60,10 @@ def all_interrupt_values(state: Any) -> list[dict[str, Any]]:
for interrupt_item in interrupts:
extracted = _extract(interrupt_item)
if extracted is not None:
values.append(extracted)
entries.append(extracted)
if saw_task_interrupt:
return values
return entries
try:
state_interrupts = getattr(state, "interrupts", ()) or ()
@ -60,5 +72,48 @@ def all_interrupt_values(state: Any) -> list[dict[str, Any]]:
for interrupt_item in state_interrupts:
extracted = _extract(interrupt_item)
if extracted is not None:
values.append(extracted)
return values
entries.append(extracted)
return entries
def all_interrupt_values(state: Any) -> list[dict[str, Any]]:
"""Interrupt payloads across the snapshot, in traversal order (ids dropped)."""
return [value for value, _ in all_interrupt_entries(state)]
def pending_interrupt_entries_from_writes(
pending_writes: Iterable[Any] | None,
) -> list[tuple[dict[str, Any], str | None]]:
"""``(value, interrupt_id)`` for interrupts stored in a checkpoint's writes.
Reads paused interrupts without compiling the agent graph, so the
thread-load path can surface HITL cards after a refresh ``aget_state``
(which would compute these) needs the full compiled agent and is far too
heavy for a read. Each ``pending_writes`` entry is ``(task_id, channel,
value)``; interrupt writes carry one or more langgraph ``Interrupt``
objects on the ``"__interrupt__"`` channel.
ponytail: couples to the persisted channel string. If langgraph renames it,
switch to ``graph.aget_state(config).interrupts``.
"""
entries: list[tuple[dict[str, Any], str | None]] = []
for write in pending_writes or ():
try:
_task_id, channel, value = write
except (ValueError, TypeError):
continue
if channel != _INTERRUPT_CHANNEL:
continue
items = value if isinstance(value, list | tuple) else [value]
for item in items:
interrupt_value = getattr(item, "value", None)
interrupt_id = getattr(item, "id", None)
if not isinstance(interrupt_value, dict):
continue
entries.append(
(
interrupt_value,
str(interrupt_id) if interrupt_id is not None else None,
)
)
return entries

View file

@ -1,7 +1,5 @@
"""Unified document save/update logic for file processors."""
import logging
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
@ -125,13 +123,6 @@ async def save_file_document(
except SQLAlchemyError as db_error:
await session.rollback()
if "ix_documents_content_hash" in str(db_error):
logging.warning(
"content_hash collision during commit for %s (%s). Skipping.",
file_name,
etl_service,
)
return None
raise db_error
except Exception as e:
await session.rollback()

View file

@ -0,0 +1,59 @@
"""Two git paths, identical bytes must converge into two rows.
"""
from __future__ import annotations
import pytest
from sqlalchemy import select
from app.config import config as app_config
from app.db import Document
from app.knowledge_store import KnowledgeStore
from app.knowledge_store.identities import user_identity
from app.knowledge_store.index.converge import index_changes
from app.utils.document_converters import generate_content_hash
pytestmark = pytest.mark.integration
DUPLICATE = "# Shared\n\nidentical bytes at two paths\n"
@pytest.fixture
def knowledge_root(tmp_path, monkeypatch):
monkeypatch.setattr(app_config, "KNOWLEDGE_STORE_ENABLED", True)
monkeypatch.setattr(app_config, "KNOWLEDGE_STORE_ROOT", str(tmp_path))
return tmp_path
@pytest.fixture
def store(knowledge_root, db_workspace):
return KnowledgeStore.for_workspace(db_workspace.id)
async def _commit_two_identical_files(store):
async with store.transaction(message="test", author=user_identity("1")) as tx:
tx.write("documents/a.xml", DUPLICATE.encode())
tx.write("documents/b.xml", DUPLICATE.encode())
return tx.revision
async def test_identical_content_at_two_paths_converges_to_two_rows(
store, db_session, db_workspace, patched_embed_texts, patched_chunk_text
):
await _commit_two_identical_files(store)
await index_changes(db_session, db_workspace.id)
rows = (
(
await db_session.execute(
select(Document).where(Document.workspace_id == db_workspace.id)
)
)
.scalars()
.all()
)
assert {row.path for row in rows} == {"/documents/a.xml", "/documents/b.xml"}
assert {row.content_hash for row in rows} == {
generate_content_hash(DUPLICATE, db_workspace.id)
}

View file

@ -72,6 +72,93 @@ class TestSliceDecisionsByToolCall:
assert routed == {}
class TestIdBasedRouting:
"""Decisions carrying a ``tool_call_id`` route by identity, not position."""
def test_routes_by_id_ignoring_submission_order(self):
decisions = [
{"type": "reject", "tool_call_id": "tcid-B"},
{"type": "approve", "tool_call_id": "tcid-A"},
]
pending = [("tcid-A", 1), ("tcid-B", 1)]
routed = slice_decisions_by_tool_call(decisions, pending)
assert routed == {
"tcid-A": {"decisions": [{"type": "approve", "tool_call_id": "tcid-A"}]},
"tcid-B": {"decisions": [{"type": "reject", "tool_call_id": "tcid-B"}]},
}
def test_groups_multi_action_bundle_by_id(self):
decisions = [
{"type": "approve", "tool_call_id": "tcid-B"},
{"type": "approve", "tool_call_id": "tcid-A"},
{"type": "edit", "tool_call_id": "tcid-A"},
]
pending = [("tcid-A", 2), ("tcid-B", 1)]
routed = slice_decisions_by_tool_call(decisions, pending)
assert routed == {
"tcid-A": {
"decisions": [
{"type": "approve", "tool_call_id": "tcid-A"},
{"type": "edit", "tool_call_id": "tcid-A"},
]
},
"tcid-B": {"decisions": [{"type": "approve", "tool_call_id": "tcid-B"}]},
}
def test_raises_on_unknown_id(self):
decisions = [{"type": "approve", "tool_call_id": "tcid-ghost"}]
pending = [("tcid-A", 1)]
with pytest.raises(ValueError, match="tcid-ghost|does not match"):
slice_decisions_by_tool_call(decisions, pending)
def test_raises_on_missing_id(self):
decisions = [{"type": "approve", "tool_call_id": "tcid-A"}]
pending = [("tcid-A", 1), ("tcid-B", 1)]
with pytest.raises(ValueError, match="tcid-B|does not match"):
slice_decisions_by_tool_call(decisions, pending)
def test_raises_on_per_id_count_mismatch(self):
decisions = [
{"type": "approve", "tool_call_id": "tcid-A"},
{"type": "approve", "tool_call_id": "tcid-A"},
]
pending = [("tcid-A", 1)]
with pytest.raises(ValueError, match="tcid-A|count"):
slice_decisions_by_tool_call(decisions, pending)
def test_partial_ids_fall_back_to_positional(self):
decisions = [
{"type": "approve"},
{"type": "reject", "tool_call_id": "tcid-B"},
]
pending = [("tcid-A", 1), ("tcid-B", 1)]
routed = slice_decisions_by_tool_call(decisions, pending)
assert routed == {
"tcid-A": {"decisions": [decisions[0]]},
"tcid-B": {"decisions": [decisions[1]]},
}
def test_null_ids_fall_back_to_positional(self):
decisions = [
{"type": "approve", "tool_call_id": None},
{"type": "reject", "tool_call_id": None},
]
pending = [("tcid-only", 2)]
routed = slice_decisions_by_tool_call(decisions, pending)
assert routed == {"tcid-only": {"decisions": decisions}}
def _interrupt_with(tool_call_id: str, action_count: int):
return SimpleNamespace(
id=f"i-{tool_call_id}",

View file

@ -0,0 +1,42 @@
"""``task`` returns a model-readable error when no tool runtime is injected.
A truncated/mangled tool call can reach the ``task`` closure without LangChain
injecting ``ToolRuntime`` (``runtime is None``). Both the sync and async paths
must degrade to a ToolMessage-able string instead of raising ``AttributeError``
on ``runtime.tool_call_id`` and killing the turn.
"""
from __future__ import annotations
import pytest
from langchain_core.runnables import RunnableLambda
from app.agents.chat.multi_agent_chat.main_agent.middleware.checkpointed_subagent_middleware.task_tool import (
build_task_tool_with_parent_config,
)
pytestmark = pytest.mark.unit
def _tool():
sub = RunnableLambda(lambda s: {"messages": []})
# workspace_id=None so the async path's spawn-paused check bypasses Redis.
return build_task_tool_with_parent_config(
[{"name": "alpha", "description": "alpha", "runnable": sub}],
workspace_id=None,
)
def test_sync_missing_runtime_returns_error_string() -> None:
out = _tool().func(description="x", subagent_type="alpha", runtime=None)
assert isinstance(out, str)
assert "runtime" in out.lower()
@pytest.mark.asyncio
async def test_async_missing_runtime_returns_error_string() -> None:
out = await _tool().coroutine(description="x", subagent_type="alpha", runtime=None)
assert isinstance(out, str)
assert "runtime" in out.lower()

View file

@ -0,0 +1,72 @@
"""``ToolRuntime`` must be injected into ``task`` through the real ToolNode.
Regression guard for the custom ``args_schema``: if the schema omits the
directly-injected ``runtime`` field, pydantic validation silently drops the
ToolNode-injected runtime and every ``task`` call fails with ``runtime=None``.
This drives the tool exactly as the agent does (StateGraph + ToolNode) and
asserts the subagent actually runs.
"""
from __future__ import annotations
from typing import Annotated
import pytest
from langchain_core.messages import AIMessage
from langchain_core.runnables import RunnableLambda
from langchain_core.utils.function_calling import convert_to_openai_tool
from langgraph.graph import END, START, StateGraph
from langgraph.graph.message import add_messages
from langgraph.prebuilt import ToolNode
from typing_extensions import TypedDict
from app.agents.chat.multi_agent_chat.main_agent.middleware.checkpointed_subagent_middleware.task_tool import (
build_task_tool_with_parent_config,
)
pytestmark = pytest.mark.unit
class _S(TypedDict):
messages: Annotated[list, add_messages]
def _tool():
sub = RunnableLambda(lambda s: {"messages": [AIMessage(content="KB ran.")]})
return build_task_tool_with_parent_config(
[{"name": "knowledge_base", "description": "kb", "runnable": sub}],
workspace_id=1,
)
@pytest.mark.asyncio
async def test_runtime_is_injected_and_subagent_runs() -> None:
tool = _tool()
g = StateGraph(_S)
g.add_node("tools", ToolNode([tool]))
g.add_edge(START, "tools")
g.add_edge("tools", END)
app = g.compile()
ai = AIMessage(
content="",
tool_calls=[
{
"name": "task",
"args": {"description": "x", "subagent_type": "knowledge_base"},
"id": "call_1",
"type": "tool_call",
}
],
)
out = await app.ainvoke({"messages": [ai]})
# If runtime were dropped, the tool returns the "could not read the tool
# runtime" guard string instead of the subagent's output.
assert "KB ran." in str(out["messages"][-1].content)
def test_runtime_stays_out_of_model_facing_schema() -> None:
props = convert_to_openai_tool(_tool())["function"]["parameters"]["properties"]
assert "runtime" not in props

View file

@ -0,0 +1,66 @@
"""``task.subagent_type`` is schema-constrained to the live roster.
Guards three properties: the provider-facing schema advertises the roster as an
enum, paused legacy-alias checkpoints still resolve (accepted-but-hidden), and an
off-roster name comes back as a model-readable error instead of crashing.
"""
from __future__ import annotations
import pytest
from langchain_core.runnables import RunnableLambda
from langchain_core.utils.function_calling import convert_to_openai_tool
from app.agents.chat.multi_agent_chat.main_agent.middleware.checkpointed_subagent_middleware.task_tool import (
build_task_tool_with_parent_config,
)
pytestmark = pytest.mark.unit
def _tool(names: list[str]):
sub = RunnableLambda(lambda s: {"messages": []})
return build_task_tool_with_parent_config(
[{"name": n, "description": n, "runnable": sub} for n in names]
)
def _subagent_type_enum(tool) -> list[str] | None:
prop = convert_to_openai_tool(tool)["function"]["parameters"]["properties"][
"subagent_type"
]
if "enum" in prop:
return prop["enum"]
for branch in prop.get("anyOf", []):
if "enum" in branch:
return branch["enum"]
return None
def test_subagent_type_advertises_roster_enum() -> None:
tool = _tool(["knowledge_base", "web_crawler"])
assert _subagent_type_enum(tool) == ["knowledge_base", "web_crawler"]
def test_legacy_alias_resolves_to_roster_and_stays_hidden() -> None:
tool = _tool(["mcp_discovery"])
assert tool.args_schema(subagent_type="gmail").subagent_type == "mcp_discovery"
assert "gmail" not in (_subagent_type_enum(tool) or [])
def test_off_roster_name_returns_model_readable_error() -> None:
tool = _tool(["knowledge_base"])
out = tool.invoke(
{
"name": "task",
"args": {"description": "x", "subagent_type": "does_not_exist"},
"id": "call_1",
"type": "tool_call",
}
)
assert out.status == "error"
assert "knowledge_base" in out.content

View file

@ -0,0 +1,67 @@
"""Reject must degrade to a ToolMessage the model can continue from, not raise.
``PermissionMiddleware`` used to ``raise RejectedError``/``CorrectedError`` on a
reject decision. Nothing caught them, so a user rejection surfaced as a 500
(``SERVER_ERROR``) and, for subagent-gated tools, killed the parent turn. Reject
must instead emit a denial ToolMessage (mirroring the deny path) so the AI/Tool
pairing stays valid and the model can adapt.
"""
from __future__ import annotations
import pytest
from langchain_core.messages import AIMessage, ToolMessage
from app.agents.chat.multi_agent_chat.shared.permissions.middleware import core
from app.agents.chat.multi_agent_chat.shared.permissions.middleware.core import (
PermissionMiddleware,
)
from app.agents.chat.multi_agent_chat.shared.permissions.model import Rule, Ruleset
pytestmark = pytest.mark.unit
def _mw(monkeypatch, decision: dict) -> PermissionMiddleware:
monkeypatch.setattr(core, "request_permission_decision", lambda **_kw: decision)
return PermissionMiddleware(
rulesets=[
Ruleset(
rules=[Rule(permission="edit_file", pattern="*", action="ask")],
origin="test",
)
]
)
def _state() -> dict:
ai = AIMessage(
content="",
tool_calls=[
{"name": "edit_file", "args": {"path": "/x"}, "id": "c1", "type": "tool_call"}
],
)
return {"messages": [ai]}
def test_reject_emits_toolmessage_and_drops_call(monkeypatch) -> None:
mw = _mw(monkeypatch, {"decision_type": "reject"})
update, _ = mw._process(_state(), None)
assert update is not None
tms = [m for m in update["messages"] if isinstance(m, ToolMessage)]
ai = next(m for m in update["messages"] if isinstance(m, AIMessage))
assert len(tms) == 1
assert tms[0].tool_call_id == "c1"
assert tms[0].status == "error"
assert ai.tool_calls == []
def test_reject_with_feedback_carries_feedback(monkeypatch) -> None:
mw = _mw(monkeypatch, {"decision_type": "reject", "feedback": "use the trash bin"})
update, _ = mw._process(_state(), None)
tms = [m for m in update["messages"] if isinstance(m, ToolMessage)]
assert len(tms) == 1
assert "use the trash bin" in tms[0].content

View file

@ -0,0 +1,173 @@
"""Guard auto-continuation on output-token truncation.
A tool-call-free text answer cut off by ``max_tokens`` should be re-invoked and
stitched until it finishes; tool-call truncation and clean finishes pass through.
"""
from __future__ import annotations
from dataclasses import dataclass, field, replace
from typing import Any
import pytest
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage
from app.agents.chat.multi_agent_chat.main_agent.middleware.continue_on_max_length import (
ContinueOnMaxLengthMiddleware,
)
@dataclass
class _FakeModel:
max_tokens: int | None = 24
@dataclass
class _FakeRequest:
model: _FakeModel
messages: list[BaseMessage] = field(default_factory=list)
def override(self, **overrides: Any) -> "_FakeRequest":
return replace(self, **overrides)
@dataclass
class _FakeResponse:
result: list[BaseMessage]
structured_response: Any = None
def _ai(text: str, *, out: int, finish: str | None, tool_calls=None) -> AIMessage:
return AIMessage(
content=text,
tool_calls=tool_calls or [],
usage_metadata={"input_tokens": 5, "output_tokens": out, "total_tokens": 5 + out},
response_metadata={"finish_reason": finish} if finish else {},
)
def _handler_from(queue: list[AIMessage]):
calls = {"n": 0}
async def handler(_request):
calls["n"] += 1
return _FakeResponse(result=[queue.pop(0)])
return handler, calls
def _text(resp: _FakeResponse) -> str:
return resp.result[-1].content
@pytest.mark.asyncio
async def test_stitches_continuation_and_stops_when_complete():
mw = ContinueOnMaxLengthMiddleware(max_continuations=3)
handler, calls = _handler_from(
[
_ai("The ocean is ", out=24, finish="length"),
_ai("vast and deep.", out=6, finish="stop"),
]
)
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
resp = await mw.awrap_model_call(req, handler)
assert _text(resp) == "The ocean is vast and deep."
assert calls["n"] == 2
@pytest.mark.asyncio
async def test_clean_finish_passes_through_untouched():
mw = ContinueOnMaxLengthMiddleware(max_continuations=3)
handler, calls = _handler_from([_ai("All done.", out=5, finish="stop")])
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
resp = await mw.awrap_model_call(req, handler)
assert _text(resp) == "All done."
assert calls["n"] == 1
@pytest.mark.asyncio
async def test_truncated_tool_call_is_not_continued():
mw = ContinueOnMaxLengthMiddleware(max_continuations=3)
truncated_tool = _ai(
"",
out=24,
finish="length",
tool_calls=[{"name": "search", "args": {"q": "x"}, "id": "t1"}],
)
handler, calls = _handler_from([truncated_tool])
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
resp = await mw.awrap_model_call(req, handler)
assert resp.result[-1].tool_calls
assert calls["n"] == 1
@pytest.mark.asyncio
async def test_respects_continuation_cap_when_model_keeps_truncating():
mw = ContinueOnMaxLengthMiddleware(max_continuations=2)
handler, calls = _handler_from(
[_ai(c, out=24, finish="length") for c in "abcd"]
)
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
resp = await mw.awrap_model_call(req, handler)
assert calls["n"] == 3 # 1 initial + 2 continuations, then give up
assert _text(resp) == "abc"
@pytest.mark.asyncio
async def test_successful_stitch_clears_truncation_marker():
from app.services import token_tracking_service as tt
acc = tt.start_turn()
acc.truncated = True
mw = ContinueOnMaxLengthMiddleware(max_continuations=3)
handler, _ = _handler_from(
[_ai("part ", out=24, finish="length"), _ai("whole.", out=6, finish="stop")]
)
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
await mw.awrap_model_call(req, handler)
assert acc.truncated is False
@pytest.mark.asyncio
async def test_cap_exhaustion_keeps_truncation_marker():
from app.services import token_tracking_service as tt
acc = tt.start_turn()
acc.truncated = True
mw = ContinueOnMaxLengthMiddleware(max_continuations=1)
handler, _ = _handler_from(
[_ai("a", out=24, finish="length"), _ai("b", out=24, finish="length")]
)
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
await mw.awrap_model_call(req, handler)
assert acc.truncated is True
@pytest.mark.asyncio
async def test_continuation_context_includes_partial_and_nudge():
mw = ContinueOnMaxLengthMiddleware(max_continuations=1)
seen: dict[str, Any] = {}
async def handler(request):
seen["messages"] = list(request.messages)
if len(seen["messages"]) > 1:
return _FakeResponse(result=[_ai("END", out=3, finish="stop")])
return _FakeResponse(result=[_ai("START ", out=24, finish="length")])
req = _FakeRequest(model=_FakeModel(max_tokens=24), messages=[HumanMessage("hi")])
resp = await mw.awrap_model_call(req, handler)
assert _text(resp) == "START END"
assert any(isinstance(m, AIMessage) and "START" in m.content for m in seen["messages"])

View file

@ -0,0 +1,41 @@
"""``ask_knowledge_base`` must self-correct, not crash the turn, on a bad call."""
from __future__ import annotations
import pytest
from langchain.tools import ToolRuntime
from app.agents.chat.multi_agent_chat.subagents.builtins.knowledge_base.ask_knowledge_base_tool import (
build_ask_knowledge_base_tool,
)
pytestmark = pytest.mark.unit
def _runtime(tool_call_id: str) -> ToolRuntime:
return ToolRuntime(
state={},
context=None,
config={},
stream_writer=None,
tool_call_id=tool_call_id,
store=None,
)
def test_missing_tool_call_id_returns_error_string() -> None:
tool = build_ask_knowledge_base_tool(kb_readonly=lambda: None)
result = tool.func("what is X?", _runtime(""))
assert isinstance(result, str)
assert "tool call id" in result.lower()
async def test_missing_tool_call_id_returns_error_string_async() -> None:
tool = build_ask_knowledge_base_tool(kb_readonly=lambda: None)
result = await tool.coroutine("what is X?", _runtime(""))
assert isinstance(result, str)
assert "tool call id" in result.lower()

View file

@ -2,15 +2,21 @@
from __future__ import annotations
from types import SimpleNamespace
import pytest
from langchain_core.messages import AIMessage
from app.agents.chat.multi_agent_chat.main_agent.middleware.tool_call_repair.builder import (
build_repair_mw,
)
from app.agents.chat.multi_agent_chat.main_agent.middleware.tool_call_repair.middleware import (
ToolCallNameRepairMiddleware,
)
from app.agents.chat.multi_agent_chat.main_agent.tools.invalid_tool import (
INVALID_TOOL_NAME,
)
from app.agents.chat.multi_agent_chat.shared.feature_flags import AgentFeatureFlags
pytestmark = pytest.mark.unit
@ -105,6 +111,37 @@ class TestRepair:
out = mw.after_model({"messages": []}, _FakeRuntime())
assert out is None
def test_unbound_deepagent_builtin_routes_to_invalid(self) -> None:
"""A router never binds FS builtins (ls/read_file/…); they must self-correct.
Previously the builder marked all deepagents builtins as "known", so a
``read_file`` call passed repair and then dispatch-failed. It must route
to ``invalid`` instead.
"""
tools = [
SimpleNamespace(name="search_knowledge_base"),
SimpleNamespace(name=INVALID_TOOL_NAME),
]
mw = build_repair_mw(flags=AgentFeatureFlags(), tools=tools)
msg = AIMessage(
content="",
tool_calls=[{"name": "read_file", "args": {"path": "/x"}, "id": "1"}],
)
out = mw.after_model(_make_state(msg), _FakeRuntime())
assert out is not None
assert out["messages"][0].tool_calls[0]["name"] == INVALID_TOOL_NAME
@pytest.mark.parametrize("name", ["task", "write_todos"])
def test_middleware_bound_tools_stay_known(self, name: str) -> None:
"""``task``/``write_todos`` are bound via middleware, not ``tools`` — keep them known."""
mw = build_repair_mw(
flags=AgentFeatureFlags(),
tools=[SimpleNamespace(name=INVALID_TOOL_NAME)],
)
msg = AIMessage(content="", tool_calls=[{"name": name, "args": {}, "id": "1"}])
out = mw.after_model(_make_state(msg), _FakeRuntime())
assert out is None # recognized, not rewritten to invalid
def test_runtime_context_extends_registered(self) -> None:
from types import SimpleNamespace

View file

@ -0,0 +1,71 @@
"""A thread paused for HITL must refuse a fresh turn.
The busy mutex releases on an ``interrupt()`` pause, so a paused thread reads
as idle and ``new_chat`` / ``regenerate`` would run over the paused checkpoint,
orphaning the pending approval. ``_raise_if_thread_awaiting_approval`` closes
that gap by reading the checkpoint and refusing with 409.
"""
import pytest
from fastapi import HTTPException
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.graph import END, START, StateGraph
from langgraph.types import interrupt
from typing_extensions import TypedDict
from app.routes.new_chat_routes import _raise_if_thread_awaiting_approval
class _S(TypedDict, total=False):
messages: list
async def _paused_checkpointer(thread_id: int) -> InMemorySaver:
"""Run a graph that interrupts, leaving a pending interrupt in the checkpoint."""
def node(_s):
decision = interrupt({"action_requests": [{"name": "x", "args": {}}]})
return {"messages": [decision]}
g = StateGraph(_S)
g.add_node("n", node)
g.add_edge(START, "n")
g.add_edge("n", END)
cp = InMemorySaver()
graph = g.compile(checkpointer=cp)
await graph.ainvoke(
{"messages": []}, {"configurable": {"thread_id": str(thread_id)}}
)
return cp
@pytest.mark.asyncio
async def test_paused_thread_is_refused_with_409():
cp = await _paused_checkpointer(1)
with pytest.raises(HTTPException) as exc:
await _raise_if_thread_awaiting_approval(1, cp)
assert exc.value.status_code == 409
assert exc.value.detail["errorCode"] == "THREAD_AWAITING_APPROVAL"
@pytest.mark.asyncio
async def test_clean_thread_is_allowed():
def node(_s):
return {"messages": ["done"]}
g = StateGraph(_S)
g.add_node("n", node)
g.add_edge(START, "n")
g.add_edge("n", END)
cp = InMemorySaver()
graph = g.compile(checkpointer=cp)
await graph.ainvoke({"messages": []}, {"configurable": {"thread_id": "2"}})
await _raise_if_thread_awaiting_approval(2, cp) # no raise
@pytest.mark.asyncio
async def test_thread_without_checkpoint_is_allowed():
await _raise_if_thread_awaiting_approval(999, InMemorySaver()) # no raise

View file

@ -0,0 +1,33 @@
"""Guard the output-truncation detector.
``langchain_litellm`` (0.6.4) drops ``finish_reason`` from streamed chunks, so
a token-limit cut reaches the UI silently. The LiteLLM success callback still
sees the real ``finish_reason`` and usage, so detection must honour both: the
gold ``finish_reason == "length"`` signal, and a usage>=max_tokens fallback for
paths where ``finish_reason`` is absent.
"""
from app.services.token_tracking_service import is_output_truncated
def test_finish_reason_length_is_truncated():
assert is_output_truncated("length", completion_tokens=10, max_tokens=999) is True
def test_finish_reason_stop_is_not_truncated_even_at_cap():
# An explicit non-length reason wins over the usage heuristic.
assert is_output_truncated("stop", completion_tokens=24, max_tokens=24) is False
def test_usage_fallback_when_finish_reason_missing():
assert is_output_truncated(None, completion_tokens=24, max_tokens=24) is True
assert is_output_truncated("", completion_tokens=30, max_tokens=24) is True
def test_under_cap_without_finish_reason_is_not_truncated():
assert is_output_truncated(None, completion_tokens=10, max_tokens=24) is False
def test_no_cap_configured_cannot_infer_from_usage():
assert is_output_truncated(None, completion_tokens=9999, max_tokens=None) is False
assert is_output_truncated(None, completion_tokens=9999, max_tokens=0) is False

View file

@ -0,0 +1,44 @@
"""``VercelStreamingService.format_interrupt_request`` carries ``interrupt_id`` on the wire.
Parent-side interrupts (doom-loop, permission asks) have no ``tool_call_id``; the
langgraph ``Interrupt.id`` is their only stable handle, so the frontend can only
render and resume them when it arrives on the frame.
"""
from __future__ import annotations
import json
import pytest
from app.services.new_streaming_service import VercelStreamingService
pytestmark = pytest.mark.unit
def _payload(frame: str) -> dict:
body = frame.removeprefix("data: ").removesuffix("\n\n")
return json.loads(body)["data"]
def test_interrupt_id_present_when_supplied() -> None:
frame = VercelStreamingService().format_interrupt_request(
{"type": "permission_ask", "action": {"tool": "search", "params": {}},
"context": {"permission": "doom_loop"}},
interrupt_id="int_7",
)
assert _payload(frame)["interrupt_id"] == "int_7"
def test_interrupt_id_omitted_when_absent() -> None:
frame = VercelStreamingService().format_interrupt_request(
{"action_requests": [], "review_configs": []},
)
assert "interrupt_id" not in _payload(frame)
def test_does_not_mutate_source_value() -> None:
"""Subagent payloads pass through by reference — stamping must not touch state."""
value = {"action_requests": [{"name": "x", "args": {}}], "review_configs": [{}]}
VercelStreamingService().format_interrupt_request(value, interrupt_id="int_9")
assert "interrupt_id" not in value

View file

@ -0,0 +1,141 @@
"""``build_resume_routing`` must route parent-side interrupts (doom-loop / asks).
Guards the bug where unstamped parent-graph interrupts were dropped on resume,
hanging the turn.
"""
from __future__ import annotations
from types import SimpleNamespace
import pytest
from app.tasks.chat.streaming.flows.resume_chat.resume_routing import (
build_resume_routing,
)
class _FakeAgent:
"""Minimal stand-in exposing the ``aget_state`` the router reads."""
def __init__(self, state):
self._state = state
async def aget_state(self, _config):
return self._state
def _doom_loop_interrupt(interrupt_id: str):
"""A parent-side doom-loop interrupt: no ``tool_call_id``, no ``action_requests``."""
return SimpleNamespace(
id=interrupt_id,
value={
"type": "permission_ask",
"action": {"tool": "search_run", "params": {}},
"context": {"permission": "doom_loop", "threshold": 3},
},
)
async def test_parent_side_doom_loop_interrupt_is_routable():
"""Decision routes to the doom-loop's ``Interrupt.id`` as a raw dict, not a bundle."""
decision = {"type": "reject"}
agent = _FakeAgent(
SimpleNamespace(interrupts=(_doom_loop_interrupt("i-doom"),))
)
routing = await build_resume_routing(agent, chat_id=42, decisions=[decision])
assert routing.lg_resume_map == {"i-doom": decision}
# No subagent bridge for parent-side interrupts.
assert routing.routed_resume_value == {}
def _permission_ask_interrupt(interrupt_id: str):
"""A parent-side main-agent permission ask: LC HITL bundle, still no stamp."""
return SimpleNamespace(
id=interrupt_id,
value={
"type": "permission_ask",
"action_requests": [{"name": "create_automation", "args": {}}],
},
)
async def test_parent_side_permission_ask_is_routable():
"""Main-agent permission asks (unstamped) route their single decision by id."""
decision = {"type": "approve"}
agent = _FakeAgent(
SimpleNamespace(interrupts=(_permission_ask_interrupt("i-perm"),))
)
routing = await build_resume_routing(agent, chat_id=7, decisions=[decision])
assert routing.lg_resume_map == {"i-perm": decision}
def _subagent_interrupt(interrupt_id: str, tool_call_id: str, action_count: int):
return SimpleNamespace(
id=interrupt_id,
value={
"action_requests": [{"name": "n", "args": {}}] * action_count,
"tool_call_id": tool_call_id,
},
)
async def test_subagent_path_is_unchanged():
"""Regression guard: stamped subagent interrupts still route via the bridge."""
decisions = [{"type": "approve"}, {"type": "reject"}]
agent = _FakeAgent(
SimpleNamespace(interrupts=(_subagent_interrupt("i-A", "tcid-A", 2),))
)
routing = await build_resume_routing(agent, chat_id=1, decisions=decisions)
assert routing.routed_resume_value == {"tcid-A": {"decisions": decisions}}
assert routing.lg_resume_map == {"i-A": {"decisions": decisions}}
async def test_id_stamped_decisions_route_by_identity_across_boundary():
"""Id-stamped decisions route correctly even reversed vs ``state.interrupts``."""
agent = _FakeAgent(
SimpleNamespace(
interrupts=(
_subagent_interrupt("i-A", "tcid-A", 1),
_subagent_interrupt("i-B", "tcid-B", 1),
)
)
)
decisions = [
{"type": "reject", "tool_call_id": "tcid-B"},
{"type": "approve", "tool_call_id": "tcid-A"},
]
routing = await build_resume_routing(agent, chat_id=1, decisions=decisions)
assert routing.routed_resume_value == {
"tcid-A": {"decisions": [{"type": "approve", "tool_call_id": "tcid-A"}]},
"tcid-B": {"decisions": [{"type": "reject", "tool_call_id": "tcid-B"}]},
}
assert routing.lg_resume_map == {
"i-A": {"decisions": [{"type": "approve", "tool_call_id": "tcid-A"}]},
"i-B": {"decisions": [{"type": "reject", "tool_call_id": "tcid-B"}]},
}
async def test_mixed_parent_and_subagent_pauses_fail_loud():
"""A pause holding both interrupt kinds is unsupported and must not mis-route."""
agent = _FakeAgent(
SimpleNamespace(
interrupts=(
_subagent_interrupt("i-A", "tcid-A", 1),
_doom_loop_interrupt("i-doom"),
)
)
)
with pytest.raises(ValueError, match="mixed HITL routing"):
await build_resume_routing(
agent, chat_id=1, decisions=[{"type": "approve"}, {"type": "reject"}]
)

View file

@ -21,11 +21,110 @@ from typing_extensions import TypedDict
from app.agents.chat.multi_agent_chat.main_agent.middleware.checkpointed_subagent_middleware.task_tool import (
build_task_tool_with_parent_config,
)
from types import SimpleNamespace
from app.tasks.chat.streaming.helpers.interrupt_inspector import (
all_interrupt_entries,
all_interrupt_values,
pending_interrupt_entries_from_writes,
)
class TestPendingInterruptEntriesFromWrites:
"""Read paused interrupts straight from a checkpoint's ``pending_writes``.
Powers refresh recovery: the thread-load path surfaces paused HITL cards
without compiling the full agent graph. Interrupts persist as writes to
the ``"__interrupt__"`` channel, one per paused task.
"""
def test_extracts_value_and_id(self):
from langgraph.types import Interrupt
writes = [
("task-1", "messages", ["ignored"]),
(
"task-2",
"__interrupt__",
[Interrupt(value={"tool_call_id": "tc-A"}, id="int-A")],
),
]
assert pending_interrupt_entries_from_writes(writes) == [
({"tool_call_id": "tc-A"}, "int-A")
]
def test_handles_scalar_value_not_wrapped_in_list(self):
from langgraph.types import Interrupt
writes = [("t", "__interrupt__", Interrupt(value={"a": 1}, id="i-1"))]
assert pending_interrupt_entries_from_writes(writes) == [({"a": 1}, "i-1")]
def test_skips_non_interrupt_channels_and_non_dict_values(self):
from langgraph.types import Interrupt
writes = [
("t", "messages", [Interrupt(value={"a": 1}, id="x")]),
("t", "__interrupt__", [Interrupt(value="not-a-dict", id="y")]),
]
assert pending_interrupt_entries_from_writes(writes) == []
def test_none_input_returns_empty(self):
assert pending_interrupt_entries_from_writes(None) == []
class TestAllInterruptEntries:
"""``all_interrupt_entries`` pairs each interrupt value with its ``Interrupt.id``.
The id is what lets parent-side interrupts (doom-loop, permission asks)
which carry no ``tool_call_id`` be addressed on the wire and on resume.
"""
def test_pairs_value_with_id_from_state_interrupts(self):
state = SimpleNamespace(
interrupts=(
SimpleNamespace(id="i-1", value={"context": {"permission": "doom_loop"}}),
SimpleNamespace(id="i-2", value={"tool_call_id": "tcid-A"}),
)
)
assert all_interrupt_entries(state) == [
({"context": {"permission": "doom_loop"}}, "i-1"),
({"tool_call_id": "tcid-A"}, "i-2"),
]
def test_prefers_task_bucket_interrupts(self):
state = SimpleNamespace(
tasks=(
SimpleNamespace(interrupts=(SimpleNamespace(id="t-1", value={"a": 1}),)),
),
interrupts=(SimpleNamespace(id="ignored", value={"b": 2}),),
)
assert all_interrupt_entries(state) == [({"a": 1}, "t-1")]
def test_skips_non_dict_values(self):
state = SimpleNamespace(
interrupts=(SimpleNamespace(id="i-1", value="not-a-dict"),)
)
assert all_interrupt_entries(state) == []
def test_id_missing_is_none(self):
state = SimpleNamespace(interrupts=(SimpleNamespace(value={"a": 1}),))
assert all_interrupt_entries(state) == [({"a": 1}, None)]
def test_values_helper_derives_from_entries(self):
state = SimpleNamespace(
interrupts=(SimpleNamespace(id="i-1", value={"a": 1}),)
)
assert all_interrupt_values(state) == [{"a": 1}]
class _SubState(TypedDict, total=False):
messages: list

View file

@ -65,7 +65,7 @@ import {
import { extractMentionedDocuments } from "@/lib/chat/stream-engine/helpers";
import { chatStreamStore } from "@/lib/chat/stream-engine/store";
import { useChatStream } from "@/lib/chat/stream-engine/use-chat-stream";
import type { ThreadRecord } from "@/lib/chat/thread-persistence";
import { getPendingInterrupts, type ThreadRecord } from "@/lib/chat/thread-persistence";
import {
extractUserTurnForNewChatApi,
type NewChatUserImagePayload,
@ -356,6 +356,50 @@ export default function NewChatPage() {
threadMessagesQuery.data,
]);
// Rebuild paused HITL cards after a refresh. The live overlay lives only in
// module RAM, so on reload we ask the backend for interrupts still pending
// in the checkpoint and repopulate the store (which re-pins the thread).
const reconstructedInterruptsRef = useRef<number | null>(null);
useEffect(() => {
if (!activeThreadId || isRunning || !threadMessagesQuery.data) return;
if (reconstructedInterruptsRef.current === activeThreadId) return;
if (chatStreamStore.getPendingInterrupts(activeThreadId).length > 0) return;
reconstructedInterruptsRef.current = activeThreadId;
const threadId = activeThreadId;
void getPendingInterrupts(threadId)
.then((resp) => {
if (resp.assistant_message_id == null || resp.pending_interrupts.length === 0) return;
if (chatStreamStore.getPendingInterrupts(threadId).length > 0) return;
const assistantMsgId = `msg-${resp.assistant_message_id}`;
const reconstructed = resp.pending_interrupts
.map((interruptData) => {
const interruptId = String(
interruptData.tool_call_id ?? interruptData.interrupt_id ?? ""
);
const actionRequests = Array.isArray(interruptData.action_requests)
? interruptData.action_requests
: [];
return {
interruptId,
threadId,
assistantMsgId,
interruptData,
bundleToolCallIds: actionRequests.map((_a, i) => `reconstructed-${interruptId}-${i}`),
} satisfies PendingInterruptState;
})
.filter((p) => p.interruptId);
if (reconstructed.length > 0) {
chatStreamStore.setPendingInterrupts(threadId, () => reconstructed);
}
})
.catch((err) => {
// Non-fatal: the thread still renders; the card just won't reappear.
console.error("[NewChatPage] Failed to load pending interrupts:", err);
reconstructedInterruptsRef.current = null;
});
}, [activeThreadId, isRunning, threadMessagesQuery.data]);
useEffect(() => {
const loadError = threadDetailQuery.error ?? threadMessagesQuery.error;
if (!activeThreadId || !loadError) return;
@ -628,6 +672,9 @@ export default function NewChatPage() {
const incoming = detail.decisions;
if (incoming.length === 0) return;
const tcIds = pendingInterrupts.flatMap((p) => p.bundleToolCallIds);
const parentInterruptIds = pendingInterrupts.flatMap((p) =>
p.bundleToolCallIds.map(() => p.interruptId)
);
const N = tcIds.length;
if (incoming.length !== N) {
@ -638,18 +685,19 @@ export default function NewChatPage() {
}
const byTcId = new Map<string, (typeof incoming)[number]>();
const submittedDecisions: typeof incoming = [];
const submittedDecisions: Array<(typeof incoming)[number] & { tool_call_id: string }> = [];
for (let i = 0; i < tcIds.length; i++) {
const tcId = tcIds[i];
const parentId = parentInterruptIds[i];
const decision = incoming[i];
if (tcId === undefined || decision === undefined) {
if (tcId === undefined || parentId === undefined || decision === undefined) {
toast.error(
`Cannot resume: ${incoming.length} decision(s) submitted for ${N} pending actions.`
);
return;
}
byTcId.set(tcId, decision);
submittedDecisions.push(decision);
submittedDecisions.push({ ...decision, tool_call_id: parentId });
}
const targetAssistantMsgId = pendingInterrupts[0].assistantMsgId;

View file

@ -9,6 +9,7 @@ import {
getToolComponent,
type TimelineToolProps,
} from "@/features/chat-messages/timeline/tool-registry";
import { isDoomLoopInterrupt } from "../approval-cards";
import type {
HitlDecision,
InterruptActionRequest,
@ -161,7 +162,12 @@ export const HitlApprovalCard: FC<{
const stagedDecision = decisions[currentStep];
const sliced = sliceForStep(interruptData, action, reviewConfig, stagedDecision);
const Body = getToolComponent(action.name) ?? FallbackToolBody;
// Doom-loop's ``action.name`` is the *stuck* tool, not an approval target,
// so its registered body (or ``NullTimelineBody``) would suppress the card.
// Route it through the HITL-aware fallback, which renders ``DoomLoopApproval``.
const Body = isDoomLoopInterrupt(sliced)
? FallbackToolBody
: (getToolComponent(action.name) ?? FallbackToolBody);
const bodyProps: TimelineToolProps = {
// Per-step key remounts the body on navigation so per-tool
// internal state (useHitlPhase, edit drafts) doesn't bleed

View file

@ -671,11 +671,13 @@ export async function startNewChat(ctx: EngineContext, message: AppendMessage):
: m
)
);
// ``tool_call_id`` is stamped on the backend by
// ``checkpointed_subagent_middleware``. Without it we can't
// address the paused subagent on resume — skip rather than
// fabricate a synthetic key.
const interruptId = String(interruptData.tool_call_id ?? "");
// Subagent interrupts carry ``tool_call_id``; parent-side ones
// (doom-loop, permission asks) carry only the langgraph
// ``interrupt_id``. Either addresses the pause on resume — skip
// only when neither is present.
const interruptId = String(
interruptData.tool_call_id ?? interruptData.interrupt_id ?? ""
);
if (interruptId) {
const incoming: PendingInterruptState = {
interruptId,
@ -822,6 +824,7 @@ export async function resumeChat(
type: string;
message?: string;
edited_action?: { name: string; args: Record<string, unknown> };
tool_call_id?: string;
}>
): Promise<void> {
const { workspaceId, threadId } = ctx;
@ -1018,7 +1021,9 @@ export async function resumeChat(
)
);
{
const interruptId = String(interruptData.tool_call_id ?? "");
const interruptId = String(
interruptData.tool_call_id ?? interruptData.interrupt_id ?? ""
);
if (interruptId) {
const incoming: PendingInterruptState = {
interruptId,

View file

@ -1,3 +1,4 @@
import { toast } from "sonner";
import {
addStepSeparator,
addToolCall,
@ -180,6 +181,12 @@ export function processSharedStreamEvent(
}
case "data-token-usage":
if (parsed.data.truncated) {
toast.warning("Response was cut off — the model hit its output-token limit.", {
duration: Infinity,
closeButton: true,
});
}
context.onTokenUsage?.(parsed.data);
return true;

View file

@ -644,6 +644,8 @@ export type SSEEvent =
total_tokens: number;
cost_micros?: number;
}>;
/** Some generation in the turn hit its output-token cap. */
truncated?: boolean;
};
}
| { type: "error"; message: string; errorCode?: string; diagnostic?: string };

View file

@ -202,6 +202,26 @@ export async function getThreadFull(threadId: number): Promise<ThreadRecord> {
return baseApiService.get<ThreadRecord>(`/api/v1/threads/${threadId}/full`);
}
export interface PendingInterruptsResponse {
/** The paused turn's assistant row; the reconstructed card reattaches here. */
assistant_message_id: number | null;
/** Each entry mirrors the ``data-interrupt-request`` SSE payload. */
pending_interrupts: Array<Record<string, unknown>>;
}
/**
* Fetch the thread's paused HITL interrupts from the LangGraph checkpoint.
*
* The live approval card lives only in ``chatStreamStore`` (module RAM), so a
* page refresh loses it while the backend stays paused. Called on thread load
* to rebuild the card.
*/
export async function getPendingInterrupts(threadId: number): Promise<PendingInterruptsResponse> {
return baseApiService.get<PendingInterruptsResponse>(
`/api/v1/threads/${threadId}/pending-interrupts`
);
}
/**
* Regeneration request parameters
*/