1286 lines
54 KiB
Python
1286 lines
54 KiB
Python
"""Run lifecycle service layer.
|
||
|
||
Centralizes the business logic for creating runs, formatting SSE
|
||
frames, and consuming stream bridge events. Router modules
|
||
(``thread_runs``, ``runs``) are thin HTTP handlers that delegate here.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import json
|
||
import logging
|
||
import os
|
||
import re
|
||
import time
|
||
import uuid
|
||
from collections.abc import Mapping
|
||
from typing import Any
|
||
|
||
from fastapi import HTTPException, Request
|
||
from langchain_core.messages import HumanMessage
|
||
|
||
from app.gateway.deps import get_concurrency_gate, get_run_context, get_run_manager, get_stream_bridge
|
||
from app.gateway.llmwiki_rag import (
|
||
LlmWikiRagContext,
|
||
attach_llmwiki_rag_to_graph_input,
|
||
build_llmwiki_rag_context,
|
||
llmwiki_selection_requested,
|
||
)
|
||
from app.gateway.utils import sanitize_log_param
|
||
from deerflow.config import get_app_config
|
||
from deerflow.config.agents_config import load_agent_config, validate_agent_id
|
||
from deerflow.config.system_settings import load_system_settings
|
||
from deerflow.runtime import (
|
||
END_SENTINEL,
|
||
HEARTBEAT_SENTINEL,
|
||
ConflictError,
|
||
DisconnectMode,
|
||
ModelConcurrencyLimitError,
|
||
RunManager,
|
||
RunRecord,
|
||
RunStatus,
|
||
StreamBridge,
|
||
UnsupportedStrategyError,
|
||
run_agent,
|
||
)
|
||
from deerflow.runtime.user_context import get_effective_user_id
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
RUN_SCOPE_EXTERNAL_ROOT = "external_root"
|
||
RUN_SCOPE_ORCHESTRATION_CHILD = "orchestration_child"
|
||
RUN_SCOPE_BACKGROUND_CHILD = "background_child"
|
||
_INTERNAL_RUN_SCOPES = {RUN_SCOPE_ORCHESTRATION_CHILD, RUN_SCOPE_BACKGROUND_CHILD}
|
||
_ALL_RUN_SCOPES = {RUN_SCOPE_EXTERNAL_ROOT, *_INTERNAL_RUN_SCOPES}
|
||
_INTERNAL_RUN_SCOPE_HEADER = "x-deerflow-internal-run-scope-token"
|
||
_INTERNAL_RUN_SCOPE_TOKEN = (
|
||
os.getenv("DEERFLOW_INTERNAL_RUN_SCOPE_TOKEN", "").strip()
|
||
or uuid.uuid4().hex
|
||
)
|
||
|
||
|
||
def internal_run_scope_headers() -> dict[str, str]:
|
||
"""Headers used by same-process loopback calls to mark internal child runs."""
|
||
|
||
return {_INTERNAL_RUN_SCOPE_HEADER: _INTERNAL_RUN_SCOPE_TOKEN}
|
||
|
||
|
||
def _metadata_run_scope(body: Any) -> str | None:
|
||
metadata = getattr(body, "metadata", None) or {}
|
||
if not isinstance(metadata, dict):
|
||
return None
|
||
raw = str(metadata.get("run_scope") or "").strip()
|
||
return raw or None
|
||
|
||
|
||
def _request_is_loopback(request: Request) -> bool:
|
||
client = getattr(request, "client", None)
|
||
host = str(getattr(client, "host", "") or "").strip()
|
||
return host in {"127.0.0.1", "::1", "localhost"}
|
||
|
||
|
||
def _has_internal_run_scope_header(request: Request) -> bool:
|
||
headers = getattr(request, "headers", None)
|
||
token = headers.get(_INTERNAL_RUN_SCOPE_HEADER) if headers is not None else None
|
||
return str(token or "") == _INTERNAL_RUN_SCOPE_TOKEN
|
||
|
||
|
||
def _resolve_run_scope(body: Any, request: Request, run_scope: str | None) -> str:
|
||
"""Return the trusted run scope used for concurrency accounting."""
|
||
|
||
if run_scope:
|
||
scope = str(run_scope).strip()
|
||
return scope if scope in _ALL_RUN_SCOPES else RUN_SCOPE_EXTERNAL_ROOT
|
||
|
||
scope = _metadata_run_scope(body)
|
||
if (
|
||
scope in _INTERNAL_RUN_SCOPES
|
||
and _request_is_loopback(request)
|
||
and _has_internal_run_scope_header(request)
|
||
):
|
||
return scope
|
||
if scope == RUN_SCOPE_EXTERNAL_ROOT:
|
||
return RUN_SCOPE_EXTERNAL_ROOT
|
||
return RUN_SCOPE_EXTERNAL_ROOT
|
||
|
||
|
||
async def _run_agent_with_concurrency_lease(
|
||
*,
|
||
concurrency_gate: Any | None,
|
||
lease_acquired: bool,
|
||
lease_metadata: dict[str, Any],
|
||
heartbeat_interval_seconds: int,
|
||
run_agent_kwargs: dict[str, Any],
|
||
) -> None:
|
||
"""Run the agent while refreshing and finally releasing an optional lease."""
|
||
|
||
heartbeat_task: asyncio.Task | None = None
|
||
|
||
async def _heartbeat_loop() -> None:
|
||
if concurrency_gate is None or not lease_acquired:
|
||
return
|
||
interval = max(5, int(heartbeat_interval_seconds or 20))
|
||
while True:
|
||
await asyncio.sleep(interval)
|
||
await concurrency_gate.heartbeat(**lease_metadata)
|
||
|
||
if concurrency_gate is not None and lease_acquired:
|
||
heartbeat_task = asyncio.create_task(
|
||
_heartbeat_loop(),
|
||
name=f"concurrency-heartbeat:{lease_metadata.get('run_id')}",
|
||
)
|
||
|
||
try:
|
||
await run_agent(**run_agent_kwargs)
|
||
finally:
|
||
if heartbeat_task is not None:
|
||
heartbeat_task.cancel()
|
||
try:
|
||
await heartbeat_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
if concurrency_gate is not None and lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# SSE formatting
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
async def immediate_end_sse():
|
||
"""Yield a single terminal ``end`` SSE frame and stop.
|
||
|
||
Used when a run exists in the persistent store but not in the in-memory
|
||
registry — typically after a Gateway restart, or once a finished run's
|
||
record has been cleaned up. There is no live stream to join, so instead of
|
||
returning 404 (which the SDK's ``joinStream`` surfaces as a hard error and
|
||
never clears its reconnect key for) we stream a clean end. The SDK treats
|
||
this as success, clears ``lg:stream:{threadId}``, and the conversation
|
||
falls back to history rendering — killing the repeating switch-back 404.
|
||
"""
|
||
yield format_sse("end", None)
|
||
|
||
|
||
def format_sse(event: str, data: Any, *, event_id: str | None = None) -> str:
|
||
"""Format a single SSE frame.
|
||
|
||
Field order: ``event:`` -> ``data:`` -> ``id:`` (optional) -> blank line.
|
||
This matches the LangGraph Platform wire format consumed by the
|
||
``useStream`` React hook and the Python ``langgraph-sdk`` SSE decoder.
|
||
"""
|
||
payload = json.dumps(data, default=str, ensure_ascii=False)
|
||
parts = [f"event: {event}", f"data: {payload}"]
|
||
if event_id:
|
||
parts.append(f"id: {event_id}")
|
||
parts.append("")
|
||
parts.append("")
|
||
return "\n".join(parts)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Input / config helpers
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def normalize_stream_modes(raw: list[str] | str | None) -> list[str]:
|
||
"""Normalize the stream_mode parameter to a list.
|
||
|
||
Default matches what ``useStream`` expects: values + messages-tuple.
|
||
"""
|
||
if raw is None:
|
||
return ["values"]
|
||
if isinstance(raw, str):
|
||
return [raw]
|
||
return raw if raw else ["values"]
|
||
|
||
|
||
def normalize_input(raw_input: dict[str, Any] | None) -> dict[str, Any]:
|
||
"""Convert LangGraph Platform input format to LangChain state dict."""
|
||
if raw_input is None:
|
||
return {}
|
||
messages = raw_input.get("messages")
|
||
if messages and isinstance(messages, list):
|
||
converted = []
|
||
for msg in messages:
|
||
if isinstance(msg, dict):
|
||
content = msg.get("content", "")
|
||
# Preserve additional_kwargs (e.g. uploaded `files`, `client_ts`,
|
||
# `prompt_prefix`) and optional id/name so they round-trip through
|
||
# the checkpoint to the UI and to UploadsMiddleware. Dropping them
|
||
# here previously made uploaded files invisible in the chat and
|
||
# broke the per-message ask-time timestamp.
|
||
msg_kwargs: dict[str, Any] = {"content": content}
|
||
extra = msg.get("additional_kwargs")
|
||
if isinstance(extra, dict) and extra:
|
||
msg_kwargs["additional_kwargs"] = extra
|
||
if msg.get("id"):
|
||
msg_kwargs["id"] = msg["id"]
|
||
if msg.get("name"):
|
||
msg_kwargs["name"] = msg["name"]
|
||
# TODO: handle other message types (system, ai, tool); for now all
|
||
# inbound dict messages are treated as human turns.
|
||
converted.append(HumanMessage(**msg_kwargs))
|
||
else:
|
||
converted.append(msg)
|
||
return {**raw_input, "messages": converted}
|
||
return raw_input
|
||
|
||
|
||
def _plain_fixed_question(graph_input: dict[str, Any]) -> str | None:
|
||
"""Return the latest plain-text human turn eligible for exact matching.
|
||
|
||
Messages carrying files or multimodal blocks are intentionally excluded:
|
||
their answer depends on more than the visible question text, so replaying a
|
||
global fixed answer would be surprising and unsafe.
|
||
"""
|
||
|
||
messages = graph_input.get("messages")
|
||
if not isinstance(messages, list) or not messages:
|
||
return None
|
||
message = messages[-1]
|
||
if not isinstance(message, HumanMessage):
|
||
return None
|
||
additional_kwargs = getattr(message, "additional_kwargs", {}) or {}
|
||
if any(
|
||
additional_kwargs.get(key)
|
||
for key in ("files", "uploaded_files", "attachments", "prompt_prefix")
|
||
):
|
||
return None
|
||
content = getattr(message, "content", None)
|
||
if isinstance(content, str):
|
||
return content
|
||
if not isinstance(content, list):
|
||
return None
|
||
parts: list[str] = []
|
||
for item in content:
|
||
if isinstance(item, str):
|
||
parts.append(item)
|
||
continue
|
||
if not isinstance(item, Mapping):
|
||
return None
|
||
item_type = str(item.get("type") or "text")
|
||
text = item.get("text")
|
||
if item_type not in {"text", "input_text"} or not isinstance(text, str):
|
||
return None
|
||
parts.append(text)
|
||
return "".join(parts)
|
||
|
||
|
||
async def _find_fixed_answer(
|
||
*,
|
||
body: Any,
|
||
request: Request,
|
||
graph_input: dict[str, Any],
|
||
run_scope: str,
|
||
) -> dict[str, Any] | None:
|
||
"""Look up an enabled exact match, falling back to the model on errors."""
|
||
|
||
if run_scope != RUN_SCOPE_EXTERNAL_ROOT or getattr(body, "command", None) is not None:
|
||
return None
|
||
assistant_id = str(getattr(body, "assistant_id", None) or _DEFAULT_ASSISTANT_ID)
|
||
if assistant_id in {"canvas_agent", "ai_writing"}:
|
||
return None
|
||
context = getattr(body, "context", None) or {}
|
||
if isinstance(context, Mapping) and context.get("writing_mode"):
|
||
return None
|
||
question = _plain_fixed_question(graph_input)
|
||
if question is None:
|
||
return None
|
||
store = getattr(request.app.state, "fixed_question_store", None)
|
||
if store is None:
|
||
return None
|
||
try:
|
||
return await store.find_enabled(question)
|
||
except Exception:
|
||
# Configuration storage must never make normal chat unavailable.
|
||
logger.warning("Fixed-question lookup failed; continuing with model", exc_info=True)
|
||
return None
|
||
|
||
|
||
_DEFAULT_ASSISTANT_ID = "lead_agent"
|
||
|
||
# Built-in LangGraph graph IDs registered in langgraph.json.
|
||
# These are NOT stored in the agents DB and must bypass the custom-agent access check.
|
||
_BUILTIN_GRAPH_IDS: frozenset[str] = frozenset({"lead_agent", "canvas_agent", "ai_writing"})
|
||
|
||
_RAG_CITATION_PROMPT = """当前对话已开启参考文献模式。
|
||
如果本轮调用了搜索、知识库、检索类技能或工具,请基于本轮所有检索结果整体回答。
|
||
如果你已经知道需要多个检索词,请优先在同一步里并行调用检索工具,等这些检索都返回后再回答,不要查一个检索词就先写一版答案。
|
||
你必须把同一轮内多次检索得到的资料视为同一个 referenceBatch,按系统稍后提供的合并批次从 [1] 开始连续编号,不要让每个检索词各自从 [1] 重新编号。
|
||
回答正文中的关键结论必须标注对应编号,例如 [1]、[2],多个来源可写为 [1][3][5]。
|
||
不要引用不存在的编号,不要使用技能内部自带的编号,不要把参考文献完整列表堆在正文末尾。
|
||
如果没有任何可用检索结果,请直接说明未检索到可靠参考资料。"""
|
||
|
||
|
||
def _prepend_prompt_prefix(graph_input: dict[str, Any], prefix: str) -> dict[str, Any]:
|
||
"""Return ``graph_input`` with the prefix *tagged* onto its first HumanMessage.
|
||
|
||
The prefix is stored in ``additional_kwargs['prompt_prefix']`` rather than
|
||
concatenated into ``content``. This keeps the persisted/displayed message
|
||
clean — the user never sees the prefix — while ``PromptPrefixMiddleware``
|
||
injects it into the content of the request actually sent to the LLM.
|
||
|
||
Leaves the input unchanged when there are no messages or the first message
|
||
is not a human turn.
|
||
"""
|
||
if not prefix:
|
||
return graph_input
|
||
messages = graph_input.get("messages")
|
||
if not isinstance(messages, list) or not messages:
|
||
return graph_input
|
||
new_messages = list(messages)
|
||
for idx, msg in enumerate(new_messages):
|
||
if isinstance(msg, HumanMessage):
|
||
extra = dict(getattr(msg, "additional_kwargs", None) or {})
|
||
extra["prompt_prefix"] = prefix
|
||
new_messages[idx] = msg.model_copy(update={"additional_kwargs": extra})
|
||
break
|
||
return {**graph_input, "messages": new_messages}
|
||
|
||
|
||
def _prompt_prefix_opted_in(body: Any) -> bool:
|
||
"""Detect whether the caller asked for the prompt prefix to be applied."""
|
||
context = getattr(body, "context", None)
|
||
if isinstance(context, Mapping) and context.get("prompt_prefix_enabled"):
|
||
return True
|
||
metadata = getattr(body, "metadata", None)
|
||
if isinstance(metadata, Mapping) and metadata.get("prompt_prefix_enabled"):
|
||
return True
|
||
return False
|
||
|
||
|
||
def _rag_mode_requested(body: Any) -> bool:
|
||
if llmwiki_selection_requested(body):
|
||
return True
|
||
context = getattr(body, "context", None)
|
||
if isinstance(context, Mapping) and context.get("rag_mode_enabled"):
|
||
return True
|
||
metadata = getattr(body, "metadata", None)
|
||
if isinstance(metadata, Mapping) and (metadata.get("rag_mode_enabled") or metadata.get("rag_mode")):
|
||
return True
|
||
return False
|
||
|
||
|
||
def _thread_rag_mode_enabled(record: Any) -> bool:
|
||
if not isinstance(record, Mapping):
|
||
return False
|
||
metadata = record.get("metadata")
|
||
return isinstance(metadata, Mapping) and bool(metadata.get("rag_mode_enabled") or metadata.get("rag_mode"))
|
||
|
||
|
||
def _append_invisible_prompt(graph_input: dict[str, Any], prompt: str) -> dict[str, Any]:
|
||
messages = graph_input.get("messages")
|
||
if not isinstance(messages, list) or not messages:
|
||
return graph_input
|
||
new_messages = list(messages)
|
||
for idx, msg in enumerate(new_messages):
|
||
if not isinstance(msg, HumanMessage):
|
||
continue
|
||
extra = dict(getattr(msg, "additional_kwargs", None) or {})
|
||
existing = str(extra.get("prompt_prefix") or "").strip()
|
||
extra["prompt_prefix"] = f"{existing}\n\n{prompt}".strip() if existing else prompt
|
||
new_messages[idx] = msg.model_copy(update={"additional_kwargs": extra})
|
||
break
|
||
return {**graph_input, "messages": new_messages}
|
||
|
||
|
||
# Whitelist of run-context keys that the langgraph-compat layer forwards from
|
||
# ``body.context`` into the run config. ``config["context"]`` exists in
|
||
# LangGraph >=0.6, but these values must be written to both ``configurable``
|
||
# (for legacy ``_get_runtime_config`` consumers) and ``context`` because
|
||
# LangGraph >=1.1.9 no longer makes ``ToolRuntime.context`` fall back to
|
||
# ``configurable`` for consumers like ``setup_agent``.
|
||
_CONTEXT_CONFIGURABLE_KEYS: frozenset[str] = frozenset(
|
||
{
|
||
"agent_id",
|
||
"model_name",
|
||
"mode",
|
||
"thinking_enabled",
|
||
# 强制关闭思考(内网 vLLM/Qwen 模型只声明 supports_thinking 时,单靠
|
||
# thinking_enabled=False 不会注入 enable_thinking=false)。透传到 run config
|
||
# 后由 lead agent 交给 create_chat_model(force_disable_thinking=...)。
|
||
"thinking_force_disabled",
|
||
"reasoning_effort",
|
||
"rag_mode_enabled",
|
||
"llmwiki_knowledge_base_ids",
|
||
"assistant_knowledge_base_ids",
|
||
"writing_mode",
|
||
"writing_artifact_path",
|
||
"canvas_artifact_title",
|
||
"canvas_artifact_type",
|
||
"is_plan_mode",
|
||
"subagent_enabled",
|
||
"max_concurrent_subagents",
|
||
"agent_name",
|
||
# 页面研报设计师由前端随消息带入当前界面明暗模式;需要进入 agent
|
||
# 提示词及 ToolRuntime.context,才能选择同色系模板。
|
||
"presentation_theme",
|
||
"is_bootstrap",
|
||
"is_writing_setup",
|
||
"is_report_structure_setup",
|
||
"valid_article_types",
|
||
"valid_retrieval_skills",
|
||
"writing_setup_quick",
|
||
"writing_sample_context",
|
||
"excluded_tools",
|
||
"skill_stop_names",
|
||
"skill_stop_message",
|
||
# 记忆开关:必须在白名单里才会从 body.context 透传到 run config,
|
||
# 否则 agent 端 cfg.get(...) 永远取默认值。圆桌等临时任务靠这两个键
|
||
# 关掉个人记忆召回 / builtin 记忆注入(见 multi_agent._run_payload)。
|
||
"memory_recall_disabled",
|
||
"memory_injection_enabled",
|
||
# 任务工作区 / 岗位会商把左侧任务卡的快照放在这里。它必须进入
|
||
# ``runtime.context``,TaskExternalContextMiddleware 才能在每次模型调用
|
||
# 前以隐藏消息注入任务名称、描述和方向;未白名单会被静默丢弃,模型只能
|
||
# 看到用户说的“这个任务”。
|
||
"task_external_context",
|
||
# 岗位会商下游席位按需读取前序席位完整交付时使用的运行期数据。
|
||
"roundtable_peer_deliveries",
|
||
# 岗位会商席位 / 方案总结:每轮 outputs 下 write_file 上限(行动规划不传此键)。
|
||
"position_artifact_cap",
|
||
# 上下文压缩开关:由聊天框下方的按钮控制(默认关闭),透传到 run config 后
|
||
# 在 _build_middlewares 里决定是否挂载 SummarizationMiddleware。
|
||
"summarization_enabled",
|
||
}
|
||
)
|
||
|
||
|
||
def merge_run_context_overrides(config: dict[str, Any], context: Mapping[str, Any] | None) -> None:
|
||
"""Merge whitelisted keys from ``body.context`` into both ``config['configurable']``
|
||
and ``config['context']`` so they are visible to legacy configurable readers and
|
||
to LangGraph ``ToolRuntime.context`` consumers (e.g. the ``setup_agent`` tool —
|
||
see issue #2677)."""
|
||
if not context:
|
||
return
|
||
configurable = config.setdefault("configurable", {})
|
||
runtime_context = config.setdefault("context", {})
|
||
for key in _CONTEXT_CONFIGURABLE_KEYS:
|
||
if key in context:
|
||
if isinstance(configurable, dict):
|
||
configurable.setdefault(key, context[key])
|
||
if isinstance(runtime_context, dict):
|
||
runtime_context.setdefault(key, context[key])
|
||
|
||
|
||
def resolve_agent_factory(assistant_id: str | None):
|
||
"""Resolve the agent factory callable from config.
|
||
|
||
Built-in graphs (``lead_agent``, ``canvas_agent``, ``ai_writing``) are dispatched
|
||
directly to their own factory. Custom agents share the ``lead_agent`` factory
|
||
with ``agent_id`` injected into configurable/context by :func:`build_run_config`.
|
||
"""
|
||
if assistant_id == "canvas_agent":
|
||
from deerflow.agents.canvas_agent.agent import make_canvas_agent
|
||
|
||
return make_canvas_agent
|
||
|
||
if assistant_id == "ai_writing":
|
||
from deerflow.agents.ai_writing.graph import make_ai_writing_graph
|
||
|
||
return make_ai_writing_graph
|
||
|
||
from deerflow.agents.lead_agent.agent import make_lead_agent
|
||
|
||
return make_lead_agent
|
||
|
||
|
||
def build_run_config(
|
||
thread_id: str,
|
||
request_config: dict[str, Any] | None,
|
||
metadata: dict[str, Any] | None,
|
||
*,
|
||
assistant_id: str | None = None,
|
||
) -> dict[str, Any]:
|
||
"""Build a RunnableConfig dict for the agent.
|
||
|
||
When *assistant_id* refers to a custom agent (anything other than
|
||
``"lead_agent"`` / ``None``), the id is forwarded as ``agent_id`` in
|
||
whichever runtime options container is active: ``context`` for
|
||
LangGraph >= 0.6.0 requests, otherwise ``configurable``.
|
||
``make_lead_agent`` reads this key to load the matching
|
||
``agents/<id>/SOUL.md`` and per-agent config — without it the agent
|
||
silently runs as the default lead agent.
|
||
|
||
This mirrors the channel manager's ``_resolve_run_params`` logic so that
|
||
the LangGraph Platform-compatible HTTP API and the IM channel path behave
|
||
identically.
|
||
"""
|
||
config: dict[str, Any] = {"recursion_limit": 250}
|
||
if request_config:
|
||
# LangGraph >= 0.6.0 introduced ``context`` as the preferred way to
|
||
# pass thread-level data and rejects requests that include both
|
||
# ``configurable`` and ``context``. If the caller already sends
|
||
# ``context``, honour it and skip our own ``configurable`` dict.
|
||
if "context" in request_config:
|
||
if "configurable" in request_config:
|
||
logger.warning(
|
||
"build_run_config: client sent both 'context' and 'configurable'; preferring 'context' (LangGraph >= 0.6.0). thread_id=%s, caller_configurable keys=%s",
|
||
thread_id,
|
||
list(request_config.get("configurable", {}).keys()),
|
||
)
|
||
context_value = request_config["context"]
|
||
if context_value is None:
|
||
context = {}
|
||
elif isinstance(context_value, Mapping):
|
||
context = dict(context_value)
|
||
else:
|
||
raise ValueError("request config 'context' must be a mapping or null.")
|
||
config["context"] = context
|
||
else:
|
||
configurable = {"thread_id": thread_id}
|
||
configurable.update(request_config.get("configurable", {}))
|
||
config["configurable"] = configurable
|
||
for k, v in request_config.items():
|
||
if k not in ("configurable", "context"):
|
||
config[k] = v
|
||
else:
|
||
config["configurable"] = {"thread_id": thread_id}
|
||
|
||
# Inject custom agent id when the caller specified a non-default assistant.
|
||
# Built-in graphs (lead_agent, canvas_agent) are dispatched via their own
|
||
# factory and must NOT receive an agent_id override.
|
||
if assistant_id and assistant_id not in _BUILTIN_GRAPH_IDS:
|
||
normalized = assistant_id.strip()
|
||
if not normalized or not re.fullmatch(r"[A-Za-z0-9_-]+", normalized):
|
||
raise ValueError(f"Invalid assistant_id {assistant_id!r}: must contain only letters, digits, hyphens, and underscores.")
|
||
if "configurable" in config:
|
||
target = config["configurable"]
|
||
elif "context" in config:
|
||
target = config["context"]
|
||
else:
|
||
target = config.setdefault("configurable", {})
|
||
if target is not None and "agent_id" not in target:
|
||
target["agent_id"] = normalized
|
||
if metadata:
|
||
config.setdefault("metadata", {}).update(metadata)
|
||
return config
|
||
|
||
|
||
def _collect_requested_agent_ids(body: Any) -> set[str]:
|
||
"""Collect agent ids from assistant_id/config/context for access checks."""
|
||
agent_ids: set[str] = set()
|
||
if body.assistant_id and body.assistant_id not in _BUILTIN_GRAPH_IDS:
|
||
agent_ids.add(str(body.assistant_id))
|
||
|
||
def collect(container: Mapping[str, Any] | None) -> None:
|
||
if not container:
|
||
return
|
||
for key in ("agent_id", "agent_name"):
|
||
value = container.get(key)
|
||
if value:
|
||
agent_ids.add(str(value))
|
||
|
||
context = getattr(body, "context", None)
|
||
if isinstance(context, Mapping):
|
||
collect(context)
|
||
|
||
config = body.config or {}
|
||
if isinstance(config, Mapping):
|
||
nested_context = config.get("context")
|
||
nested_configurable = config.get("configurable")
|
||
if isinstance(nested_context, Mapping):
|
||
collect(nested_context)
|
||
if isinstance(nested_configurable, Mapping):
|
||
collect(nested_configurable)
|
||
|
||
return agent_ids
|
||
|
||
|
||
def _extract_requested_model_name(body: Any) -> str | None:
|
||
"""Best-effort read of the model override requested for this run."""
|
||
context = getattr(body, "context", None)
|
||
if isinstance(context, Mapping):
|
||
for key in ("model_name", "model"):
|
||
value = str(context.get(key) or "").strip()
|
||
if value:
|
||
return value
|
||
|
||
config = getattr(body, "config", None)
|
||
if isinstance(config, Mapping):
|
||
for container_key in ("context", "configurable"):
|
||
container = config.get(container_key)
|
||
if not isinstance(container, Mapping):
|
||
continue
|
||
for key in ("model_name", "model"):
|
||
value = str(container.get(key) or "").strip()
|
||
if value:
|
||
return value
|
||
return None
|
||
|
||
|
||
def _extract_effective_agent_id(body: Any) -> str | None:
|
||
"""Best-effort read of the custom agent id that the lead graph will run."""
|
||
candidates: list[str] = []
|
||
|
||
def collect(container: Mapping[str, Any] | None) -> None:
|
||
if not container:
|
||
return
|
||
for key in ("agent_id", "agent_name"):
|
||
value = str(container.get(key) or "").strip()
|
||
if value:
|
||
candidates.append(value)
|
||
|
||
context = getattr(body, "context", None)
|
||
if isinstance(context, Mapping):
|
||
collect(context)
|
||
|
||
config = getattr(body, "config", None)
|
||
if isinstance(config, Mapping):
|
||
nested_context = config.get("context")
|
||
nested_configurable = config.get("configurable")
|
||
if isinstance(nested_context, Mapping):
|
||
collect(nested_context)
|
||
if isinstance(nested_configurable, Mapping):
|
||
collect(nested_configurable)
|
||
|
||
assistant_id = str(getattr(body, "assistant_id", "") or "").strip()
|
||
if assistant_id and assistant_id not in _BUILTIN_GRAPH_IDS:
|
||
candidates.append(assistant_id)
|
||
|
||
for agent_id in candidates:
|
||
try:
|
||
normalized = validate_agent_id(agent_id)
|
||
except ValueError:
|
||
logger.debug("Ignoring invalid agent id during model resolution: %s", agent_id)
|
||
continue
|
||
if normalized:
|
||
return normalized
|
||
return None
|
||
|
||
|
||
def _resolve_effective_run_model_name(body: Any) -> str:
|
||
"""Resolve the model that this run will effectively use for concurrency limiting."""
|
||
app_config = get_app_config()
|
||
default_model_name = app_config.models[0].name if app_config.models else None
|
||
if not default_model_name:
|
||
raise ValueError("No chat models are configured. Please configure at least one model in config.yaml.")
|
||
|
||
requested_model_name = _extract_requested_model_name(body)
|
||
if requested_model_name:
|
||
if app_config.get_model_config(requested_model_name):
|
||
return requested_model_name
|
||
logger.warning(
|
||
"Requested model '%s' not found; falling back to default model '%s'.",
|
||
requested_model_name,
|
||
default_model_name,
|
||
)
|
||
return default_model_name
|
||
|
||
agent_id = _extract_effective_agent_id(body)
|
||
if agent_id:
|
||
try:
|
||
agent_config = load_agent_config(agent_id)
|
||
except FileNotFoundError:
|
||
agent_config = None
|
||
except Exception:
|
||
logger.debug("Failed to load agent config for model resolution: %s", agent_id, exc_info=True)
|
||
agent_config = None
|
||
agent_model_name = str(getattr(agent_config, "model", "") or "").strip()
|
||
if agent_model_name:
|
||
if app_config.get_model_config(agent_model_name):
|
||
return agent_model_name
|
||
logger.warning(
|
||
"Agent model '%s' for agent '%s' not found; falling back to default model '%s'.",
|
||
agent_model_name,
|
||
agent_id,
|
||
default_model_name,
|
||
)
|
||
|
||
return default_model_name
|
||
|
||
|
||
def _resolve_effective_model_name(body: Any, app_config: Any | None = None) -> str:
|
||
"""Compatibility wrapper for resolving the effective run model."""
|
||
if app_config is None:
|
||
return _resolve_effective_run_model_name(body)
|
||
|
||
default_model_name = app_config.models[0].name if app_config.models else None
|
||
if not default_model_name:
|
||
raise ValueError("No chat models are configured. Please configure at least one model in config.yaml.")
|
||
|
||
requested_model_name = _extract_requested_model_name(body)
|
||
if requested_model_name and app_config.get_model_config(requested_model_name):
|
||
return requested_model_name
|
||
|
||
agent_id = _extract_effective_agent_id(body)
|
||
if agent_id:
|
||
try:
|
||
agent_config = load_agent_config(agent_id)
|
||
except Exception:
|
||
agent_config = None
|
||
agent_model_name = str(getattr(agent_config, "model", "") or "").strip()
|
||
if agent_model_name and app_config.get_model_config(agent_model_name):
|
||
return agent_model_name
|
||
|
||
return default_model_name
|
||
|
||
|
||
def _resolve_model_concurrency_limit(model_name: str, app_config: Any) -> int | None:
|
||
"""Read legacy per-model max-concurrency values from model config extras."""
|
||
model_config = app_config.get_model_config(model_name)
|
||
extras = getattr(model_config, "model_extra", None) or {}
|
||
for key in ("max_concurrency", "model_concurrency_limit"):
|
||
raw = extras.get(key)
|
||
if raw is None:
|
||
continue
|
||
try:
|
||
limit = int(raw)
|
||
except (TypeError, ValueError):
|
||
continue
|
||
if limit > 0:
|
||
return limit
|
||
return None
|
||
|
||
|
||
def _get_model_concurrency_limit(model_name: str) -> int | None:
|
||
"""Return the configured limit for a model, or ``None`` when limiting is off."""
|
||
settings = load_system_settings().model_concurrency
|
||
if not settings.enabled:
|
||
return None
|
||
|
||
limit = settings.model_limits.get(model_name, settings.default_limit)
|
||
try:
|
||
limit = int(limit)
|
||
except (TypeError, ValueError):
|
||
limit = int(settings.default_limit)
|
||
return max(1, limit)
|
||
|
||
|
||
def _get_user_model_concurrency_limit(model_name: str) -> int | None:
|
||
"""Return the configured per-user limit for a model, or ``None`` when off."""
|
||
settings = load_system_settings().model_concurrency
|
||
if not settings.enabled:
|
||
return None
|
||
|
||
raw_limit = settings.user_model_limits.get(
|
||
model_name,
|
||
getattr(settings, "user_default_limit", 0) or 0,
|
||
)
|
||
try:
|
||
limit = int(raw_limit)
|
||
except (TypeError, ValueError):
|
||
limit = int(getattr(settings, "user_default_limit", 0) or 0)
|
||
if limit < 1:
|
||
return None
|
||
return max(1, limit)
|
||
|
||
|
||
async def _sync_legacy_agent_records(store) -> None:
|
||
"""Import existing folder-backed agents as built-in records.
|
||
|
||
Heavy: ``list_custom_agents()`` scans ``.deer-flow/agents/`` and reads
|
||
``config.yaml`` + ``SOUL.md`` per directory, then ``store.ensure_builtin``
|
||
does one DB round-trip per agent. On a machine with N agent dirs this is
|
||
O(N) file IO + O(N) DB calls — historically the worst per-request hotspot
|
||
in ``start_run`` (observed 30s+ on machines with 100+ residual coordinator
|
||
shells + slow MySQL). **Never call this directly from request handlers** —
|
||
use :func:`_maybe_sync_legacy_agent_records` which applies the same TTL
|
||
gating as ``routers/agents._maybe_sync_legacy_agents``.
|
||
"""
|
||
from deerflow.config.agents_config import list_custom_agents
|
||
|
||
for agent_cfg in list_custom_agents():
|
||
agent_id = agent_cfg.id or agent_cfg.name
|
||
if not agent_id:
|
||
continue
|
||
await store.ensure_builtin(
|
||
{
|
||
"id": agent_id,
|
||
"name": agent_cfg.name,
|
||
"description": agent_cfg.description or "",
|
||
"published": True,
|
||
}
|
||
)
|
||
|
||
|
||
# Throttle per-request legacy-sync. Matches ``routers/agents._LEGACY_SYNC_TTL_SECONDS``
|
||
# semantics — the real sync runs once at startup (lifespan); this is a safety
|
||
# net for environments that bypass the lifespan hook or for freshly dropped
|
||
# agent directories. Lifespan registers itself via the same state by calling
|
||
# the agents router's ``_mark_legacy_sync_fresh()`` after its boot-time sync,
|
||
# so this module's TTL state is independent (a separate gate kept here so we
|
||
# don't import from ``routers/agents`` and create a circular import).
|
||
_LEGACY_SYNC_TTL_SECONDS = 3600.0
|
||
_legacy_sync_state: dict[str, Any] = {"last_at": 0.0, "lock": None}
|
||
|
||
|
||
# Anything slower than this triggers the WARNING-level breakdown in
|
||
# ``start_run``. 500ms gives plenty of headroom for cold caches + a few DB
|
||
# round-trips on healthy systems while still catching the pathological
|
||
# multi-second cases that brought us here.
|
||
_START_RUN_SLOW_THRESHOLD_SECONDS = 0.5
|
||
|
||
|
||
def _mark_legacy_sync_fresh() -> None:
|
||
"""Tell this module's TTL gate that a sync just completed elsewhere.
|
||
|
||
Called by the Gateway lifespan hook after it runs the boot-time legacy
|
||
agent sync (``routers/agents._sync_legacy_agents``). Without this, the
|
||
first request after startup would race past the TTL gate here and trigger
|
||
a redundant rescan even though the data is already fresh.
|
||
"""
|
||
_legacy_sync_state["last_at"] = time.monotonic()
|
||
|
||
|
||
async def _maybe_sync_legacy_agent_records(store) -> None:
|
||
"""TTL-gated wrapper around :func:`_sync_legacy_agent_records`.
|
||
|
||
Steady-state behavior: returns in microseconds inside the TTL window — no
|
||
filesystem walk, no DB round-trips. Outside the TTL window we serialize
|
||
callers under an :class:`asyncio.Lock` so a burst of concurrent
|
||
``/runs/stream`` requests fans into a single rescan instead of N parallel
|
||
fan-outs. Pattern mirrors ``routers/agents._maybe_sync_legacy_agents``.
|
||
"""
|
||
now = time.monotonic()
|
||
if now - float(_legacy_sync_state["last_at"]) < _LEGACY_SYNC_TTL_SECONDS:
|
||
return
|
||
lock = _legacy_sync_state["lock"]
|
||
if lock is None:
|
||
lock = asyncio.Lock()
|
||
_legacy_sync_state["lock"] = lock
|
||
async with lock:
|
||
# Re-check under the lock — a sibling coroutine may have just refreshed.
|
||
if time.monotonic() - float(_legacy_sync_state["last_at"]) < _LEGACY_SYNC_TTL_SECONDS:
|
||
return
|
||
t_sync_start = time.perf_counter()
|
||
await _sync_legacy_agent_records(store)
|
||
elapsed = time.perf_counter() - t_sync_start
|
||
_legacy_sync_state["last_at"] = time.monotonic()
|
||
# When this fires it's a "first request of the hour" warmup — log it
|
||
# so operators can see the spike isn't an isolated bug. Anything
|
||
# over 2s usually means a lot of agent directories + slow IO/DB.
|
||
logger.info(
|
||
"[start-run-timing] legacy-agent sync warmed cache in %.3fs (next sync in %.0fs)",
|
||
elapsed,
|
||
_LEGACY_SYNC_TTL_SECONDS,
|
||
)
|
||
|
||
|
||
async def _enforce_agent_access(body: Any, request: Request) -> None:
|
||
agent_ids = _collect_requested_agent_ids(body)
|
||
if not agent_ids:
|
||
return
|
||
|
||
from app.gateway.deps import get_agent_store
|
||
from deerflow.config.agents_config import validate_agent_id
|
||
from deerflow.runtime.user_context import get_effective_user_id
|
||
|
||
store = get_agent_store(request)
|
||
await _maybe_sync_legacy_agent_records(store)
|
||
user_id = get_effective_user_id()
|
||
# Admins see every agent in the admin console listing (?admin=true), so
|
||
# they may also start a run against any of them.
|
||
# ``request`` may be an in-process virtual request (SimpleNamespace) from
|
||
# the roundtable job gateway / scheduled-task runtime — those carry no
|
||
# ``state`` attribute, so the whole chain must be getattr-guarded
|
||
# (a bare ``request.state`` raises AttributeError and kills the run).
|
||
is_admin = getattr(getattr(getattr(request, "state", None), "user", None), "system_role", None) == "admin"
|
||
for agent_id in agent_ids:
|
||
try:
|
||
normalized = validate_agent_id(agent_id)
|
||
except ValueError as exc:
|
||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||
if normalized is None:
|
||
continue
|
||
if await store.get_visible(normalized, user_id) is not None:
|
||
continue
|
||
if is_admin and await store.get_any(normalized) is not None:
|
||
continue
|
||
raise HTTPException(status_code=404, detail=f"Agent {normalized} not found")
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Run lifecycle
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
async def start_run(
|
||
body: Any,
|
||
thread_id: str,
|
||
request: Request,
|
||
*,
|
||
enforce_access: bool = True,
|
||
run_scope: str | None = None,
|
||
) -> RunRecord:
|
||
"""Create a RunRecord and launch the background agent task.
|
||
|
||
Parameters
|
||
----------
|
||
body : RunCreateRequest
|
||
The validated request body (typed as Any to avoid circular import
|
||
with the router module that defines the Pydantic model).
|
||
thread_id : str
|
||
Target thread.
|
||
request : Request
|
||
FastAPI request — used to retrieve singletons from ``app.state``.
|
||
|
||
Timing observability
|
||
--------------------
|
||
Each phase records ``perf_counter`` deltas. When the total wallclock from
|
||
function entry to ``run_agent`` task creation exceeds
|
||
``_START_RUN_SLOW_THRESHOLD_SECONDS`` we emit a ``[start-run-timing]``
|
||
breakdown at WARNING level. Stays at DEBUG when fast, so steady-state
|
||
logs are not polluted. Investigate spikes by grep'ing
|
||
``[start-run-timing]`` together with ``[upstream-timing]`` (multi-agent
|
||
router) and ``[leader-timing] / [special-timing]`` (圆桌 leader/special).
|
||
"""
|
||
t_entry = time.perf_counter()
|
||
bridge = get_stream_bridge(request)
|
||
run_mgr = get_run_manager(request)
|
||
run_ctx = get_run_context(request)
|
||
concurrency_gate = get_concurrency_gate(request)
|
||
|
||
disconnect = DisconnectMode.cancel if body.on_disconnect == "cancel" else DisconnectMode.continue_
|
||
|
||
t_pre_enforce = time.perf_counter()
|
||
# ``enforce_access=False`` is for intentionally open/unauthenticated entry
|
||
# points (e.g. /api/open/3qfx/ask) that resolve the agent themselves from an
|
||
# admin-configured mapping rather than from a per-user-visible assistant_id.
|
||
# Those run under the "default" bucket, where a user-owned unpublished agent
|
||
# would otherwise 404 the access check.
|
||
if enforce_access:
|
||
await _enforce_agent_access(body, request)
|
||
t_post_enforce = time.perf_counter()
|
||
|
||
effective_user_id = get_effective_user_id()
|
||
effective_run_scope = _resolve_run_scope(body, request, run_scope)
|
||
graph_input = normalize_input(body.input)
|
||
fixed_answer = await _find_fixed_answer(
|
||
body=body,
|
||
request=request,
|
||
graph_input=graph_input,
|
||
run_scope=effective_run_scope,
|
||
)
|
||
if fixed_answer is not None:
|
||
effective_model_name = "fixed-answer"
|
||
model_concurrency_limit = None
|
||
user_model_concurrency_limit = None
|
||
else:
|
||
effective_model_name = _resolve_effective_run_model_name(body)
|
||
model_concurrency_limit = _get_model_concurrency_limit(effective_model_name)
|
||
user_model_concurrency_limit = (
|
||
_get_user_model_concurrency_limit(effective_model_name)
|
||
if effective_run_scope == RUN_SCOPE_EXTERNAL_ROOT
|
||
else None
|
||
)
|
||
run_id = str(uuid.uuid4())
|
||
|
||
run_metadata = dict(body.metadata or {})
|
||
run_metadata["model_name"] = effective_model_name
|
||
run_metadata["user_id"] = effective_user_id
|
||
run_metadata["run_scope"] = effective_run_scope
|
||
if fixed_answer is not None:
|
||
run_metadata["fixed_answer"] = True
|
||
run_metadata["fixed_question_id"] = fixed_answer["id"]
|
||
|
||
redis_lease_acquired = False
|
||
redis_gate_degraded = True
|
||
user_model_accounted = user_model_concurrency_limit is not None
|
||
lease_metadata = {
|
||
"run_id": run_id,
|
||
"thread_id": thread_id,
|
||
"model_name": effective_model_name,
|
||
"user_id": effective_user_id,
|
||
"user_model_accounted": user_model_accounted,
|
||
}
|
||
if concurrency_gate is not None and fixed_answer is None:
|
||
acquire = await concurrency_gate.try_acquire(
|
||
run_id=run_id,
|
||
thread_id=thread_id,
|
||
model_name=effective_model_name,
|
||
user_id=effective_user_id,
|
||
assistant_id=body.assistant_id or "",
|
||
model_limit=model_concurrency_limit,
|
||
user_model_limit=user_model_concurrency_limit,
|
||
reject_thread=body.multitask_strategy == "reject",
|
||
)
|
||
redis_lease_acquired = acquire.acquired
|
||
redis_gate_degraded = acquire.degraded
|
||
if acquire.rejected:
|
||
if acquire.reason == "thread":
|
||
raise HTTPException(
|
||
status_code=409,
|
||
detail="Current thread already has an active run. Please wait or stop it before sending another message.",
|
||
)
|
||
if acquire.reason == "user_model":
|
||
raise HTTPException(
|
||
status_code=429,
|
||
detail="This account has reached the concurrency limit for the selected model. Please try again later.",
|
||
)
|
||
raise HTTPException(status_code=429, detail="The selected model is busy. Please try again later.")
|
||
|
||
try:
|
||
record = await run_mgr.create_or_reject(
|
||
thread_id,
|
||
body.assistant_id,
|
||
on_disconnect=disconnect,
|
||
metadata=run_metadata,
|
||
kwargs={"input": body.input, "config": body.config},
|
||
multitask_strategy=body.multitask_strategy,
|
||
model_concurrency_limit=(effective_model_name, model_concurrency_limit) if redis_gate_degraded and model_concurrency_limit is not None else None,
|
||
user_model_concurrency_limit=(effective_model_name, effective_user_id, user_model_concurrency_limit) if redis_gate_degraded and user_model_concurrency_limit is not None else None,
|
||
run_id=run_id,
|
||
)
|
||
except ModelConcurrencyLimitError as exc:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
if getattr(exc, "scope", "") == "user_model":
|
||
detail = "This account has reached the concurrency limit for the selected model. Please try again later."
|
||
else:
|
||
detail = "The selected model is busy. Please try again later."
|
||
raise HTTPException(status_code=429, detail=detail) from exc
|
||
except ConflictError as exc:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
# 同一会话上一轮回答还没结束,又提交了新问题 → 拒绝。给用户中文提示,
|
||
# 而不是把内部英文「Thread ... already has an active run」直接抛到前端。
|
||
raise HTTPException(
|
||
status_code=409,
|
||
detail="Current thread already has an active run. Please wait or stop it before sending another message.",
|
||
) from exc
|
||
except UnsupportedStrategyError as exc:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
raise HTTPException(status_code=501, detail=str(exc)) from exc
|
||
except Exception:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
raise
|
||
t_post_create = time.perf_counter()
|
||
|
||
# Upsert thread metadata so the thread appears in /threads/search,
|
||
# even for threads that were never explicitly created via POST /threads
|
||
# (e.g. stateless runs).
|
||
requested_rag_mode = _rag_mode_requested(body)
|
||
effective_rag_mode = requested_rag_mode
|
||
try:
|
||
existing = await run_ctx.thread_store.get(thread_id)
|
||
effective_rag_mode = requested_rag_mode or _thread_rag_mode_enabled(existing)
|
||
if existing is None:
|
||
metadata = dict(body.metadata or {})
|
||
if effective_rag_mode:
|
||
metadata["rag_mode_enabled"] = True
|
||
metadata["rag_mode"] = True
|
||
await run_ctx.thread_store.create(
|
||
thread_id,
|
||
assistant_id=body.assistant_id,
|
||
metadata=metadata,
|
||
)
|
||
else:
|
||
await run_ctx.thread_store.update_status(thread_id, "running")
|
||
if requested_rag_mode and not _thread_rag_mode_enabled(existing):
|
||
await run_ctx.thread_store.update_metadata(
|
||
thread_id,
|
||
{"rag_mode_enabled": True, "rag_mode": True},
|
||
)
|
||
effective_rag_mode = True
|
||
# SDK 的 ``client.threads.create()`` 不传 assistant_id —— 先建出来空的 thread,
|
||
# 真正提交 run 时才知道用哪个 assistant。这里回填一次,避免历史列表 / 按
|
||
# assistant 筛选时拿不到关联。
|
||
if not existing.get("assistant_id") and body.assistant_id:
|
||
await run_ctx.thread_store.update_assistant_id(thread_id, body.assistant_id)
|
||
except Exception:
|
||
logger.warning("Failed to upsert thread_meta for %s (non-fatal)", sanitize_log_param(thread_id))
|
||
t_post_thread_upsert = time.perf_counter()
|
||
|
||
try:
|
||
if fixed_answer is not None:
|
||
from deerflow.agents.fixed_answer_agent import make_fixed_answer_agent
|
||
|
||
def agent_factory(config): # noqa: ANN001, ARG001 - worker factory contract
|
||
return make_fixed_answer_agent(
|
||
question=str(fixed_answer["question"]),
|
||
answer=str(fixed_answer["answer"]),
|
||
fixed_question_id=str(fixed_answer["id"]),
|
||
tokens_per_second=int(fixed_answer.get("tokens_per_second", 100)),
|
||
)
|
||
else:
|
||
agent_factory = resolve_agent_factory(body.assistant_id)
|
||
except Exception:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
raise
|
||
|
||
# Admin-configurable prompt prefix: when the caller opted in via
|
||
# context.prompt_prefix_enabled, tag the human message with the configured
|
||
# prefix (stored in additional_kwargs, not concatenated into content, so the
|
||
# user never sees it — PromptPrefixMiddleware injects it for the LLM only).
|
||
# The frontend only sends the flag on the opening turn, so this naturally
|
||
# applies once. We intentionally do NOT gate on thread-row existence: the
|
||
# thread_meta row is frequently pre-created, which previously skipped the
|
||
# prefix for genuine first turns.
|
||
if fixed_answer is None and _prompt_prefix_opted_in(body):
|
||
prefix_settings = load_system_settings().prompt_prefix
|
||
if prefix_settings.enabled and prefix_settings.prompt_prefix.strip():
|
||
graph_input = _prepend_prompt_prefix(graph_input, prefix_settings.prompt_prefix.strip())
|
||
llmwiki_rag = (
|
||
LlmWikiRagContext([], [], 0, "", [], "")
|
||
if fixed_answer is not None
|
||
else await build_llmwiki_rag_context(
|
||
request=request,
|
||
body=body,
|
||
graph_input=graph_input,
|
||
)
|
||
)
|
||
if llmwiki_rag.prompt or llmwiki_rag.sources:
|
||
graph_input = attach_llmwiki_rag_to_graph_input(graph_input, llmwiki_rag)
|
||
effective_rag_mode = True
|
||
run_metadata["llmwiki_knowledge_base_ids"] = llmwiki_rag.knowledge_base_ids
|
||
run_metadata["assistant_knowledge_base_ids"] = llmwiki_rag.assistant_knowledge_base_ids
|
||
run_metadata["llmwiki_rag_source_count"] = llmwiki_rag.source_count
|
||
try:
|
||
await run_ctx.thread_store.update_metadata(
|
||
thread_id,
|
||
{
|
||
"llmwiki_knowledge_base_ids": llmwiki_rag.knowledge_base_ids,
|
||
"assistant_knowledge_base_ids": llmwiki_rag.assistant_knowledge_base_ids,
|
||
"rag_mode_enabled": True,
|
||
"rag_mode": True,
|
||
},
|
||
)
|
||
except Exception:
|
||
logger.debug("Failed to persist LLMWiki RAG metadata for %s", sanitize_log_param(thread_id), exc_info=True)
|
||
if effective_rag_mode:
|
||
graph_input = _append_invisible_prompt(graph_input, _RAG_CITATION_PROMPT)
|
||
config = build_run_config(thread_id, body.config, body.metadata, assistant_id=body.assistant_id)
|
||
|
||
# Merge DeerFlow-specific context overrides into both ``configurable`` and ``context``.
|
||
# The ``context`` field is a custom extension for the langgraph-compat layer
|
||
# that carries agent configuration (model_name, thinking_enabled, etc.).
|
||
# Only agent-relevant keys are forwarded; unknown keys (e.g. thread_id) are ignored.
|
||
body_context = dict(getattr(body, "context", None) or {})
|
||
if effective_rag_mode:
|
||
body_context["rag_mode_enabled"] = True
|
||
if "llmwiki_knowledge_base_ids" not in body_context and llmwiki_rag.knowledge_base_ids:
|
||
body_context["llmwiki_knowledge_base_ids"] = llmwiki_rag.knowledge_base_ids
|
||
if "assistant_knowledge_base_ids" not in body_context and llmwiki_rag.assistant_knowledge_base_ids:
|
||
body_context["assistant_knowledge_base_ids"] = llmwiki_rag.assistant_knowledge_base_ids
|
||
if not llmwiki_rag.knowledge_base_ids:
|
||
existing_excluded = list(body_context.get("excluded_tools") or getattr(body, "excluded_tools", None) or [])
|
||
if "llmwiki_search" not in existing_excluded:
|
||
body_context["excluded_tools"] = [*existing_excluded, "llmwiki_search"]
|
||
# Promote whitelisted top-level body fields into the context map so
|
||
# callers can pass them either nested under ``context`` or as flat
|
||
# top-level fields (``excluded_tools``, ``skill_stop_names``, …).
|
||
for top_level_key in ("excluded_tools", "skill_stop_names", "skill_stop_message"):
|
||
value = getattr(body, top_level_key, None)
|
||
if value is not None and top_level_key not in body_context:
|
||
body_context[top_level_key] = value
|
||
merge_run_context_overrides(config, body_context or None)
|
||
|
||
stream_modes = normalize_stream_modes(body.stream_mode)
|
||
|
||
t_pre_task = time.perf_counter()
|
||
heartbeat_interval_seconds = 20
|
||
if concurrency_gate is not None:
|
||
heartbeat_interval_seconds = int(
|
||
getattr(getattr(concurrency_gate, "config", None), "heartbeat_interval_seconds", 20) or 20
|
||
)
|
||
try:
|
||
task = asyncio.create_task(
|
||
_run_agent_with_concurrency_lease(
|
||
concurrency_gate=concurrency_gate,
|
||
lease_acquired=redis_lease_acquired,
|
||
lease_metadata=lease_metadata,
|
||
heartbeat_interval_seconds=heartbeat_interval_seconds,
|
||
run_agent_kwargs={
|
||
"bridge": bridge,
|
||
"run_manager": run_mgr,
|
||
"record": record,
|
||
"ctx": run_ctx,
|
||
"agent_factory": agent_factory,
|
||
"graph_input": graph_input,
|
||
"config": config,
|
||
"stream_modes": stream_modes,
|
||
"stream_subgraphs": body.stream_subgraphs,
|
||
"interrupt_before": body.interrupt_before,
|
||
"interrupt_after": body.interrupt_after,
|
||
"command": getattr(body, "command", None),
|
||
},
|
||
)
|
||
)
|
||
except Exception:
|
||
if concurrency_gate is not None and redis_lease_acquired:
|
||
await concurrency_gate.release(**lease_metadata)
|
||
raise
|
||
record.task = task
|
||
|
||
# Title sync is handled by worker.py's finally block which reads the
|
||
# title from the checkpoint and calls thread_store.update_display_name
|
||
# after the run completes.
|
||
|
||
total = t_pre_task - t_entry
|
||
if total >= _START_RUN_SLOW_THRESHOLD_SECONDS:
|
||
# Slow path: emit a phase breakdown so on-call can identify the
|
||
# culprit without re-running with extra instrumentation. The most
|
||
# common cause is a stale ``_maybe_sync_legacy_agent_records`` TTL
|
||
# window firing on a machine with hundreds of agent directories.
|
||
logger.warning(
|
||
"[start-run-timing] thread=%s total=%.3fs enforce_access=%.3fs run_create=%.3fs thread_upsert=%.3fs build_payload=%.3fs",
|
||
sanitize_log_param(thread_id),
|
||
total,
|
||
t_post_enforce - t_pre_enforce,
|
||
t_post_create - t_post_enforce,
|
||
t_post_thread_upsert - t_post_create,
|
||
t_pre_task - t_post_thread_upsert,
|
||
)
|
||
else:
|
||
logger.debug(
|
||
"[start-run-timing] thread=%s total=%.3fs enforce=%.3fs create=%.3fs upsert=%.3fs build=%.3fs",
|
||
sanitize_log_param(thread_id),
|
||
total,
|
||
t_post_enforce - t_pre_enforce,
|
||
t_post_create - t_post_enforce,
|
||
t_post_thread_upsert - t_post_create,
|
||
t_pre_task - t_post_thread_upsert,
|
||
)
|
||
|
||
return record
|
||
|
||
|
||
async def sse_consumer(
|
||
bridge: StreamBridge,
|
||
record: RunRecord,
|
||
request: Request,
|
||
run_mgr: RunManager,
|
||
):
|
||
"""Async generator that yields SSE frames from the bridge.
|
||
|
||
The ``finally`` block implements ``on_disconnect`` semantics:
|
||
- ``cancel``: abort the background task on client disconnect.
|
||
- ``continue``: let the task run so the client can ``joinStream``.
|
||
|
||
``is_disconnected()`` is only consulted in cancel mode. In continue
|
||
(resumable) mode a false-positive disconnect would *cleanly* end the
|
||
HTTP stream; the LangGraph SDK treats that as success, clears
|
||
``lg:stream:{threadId}``, and flips ``isLoading`` off — while the
|
||
backend run keeps the thread occupied, so the next send 409s.
|
||
A truly gone client fails on the next ``yield`` instead.
|
||
"""
|
||
last_event_id = request.headers.get("Last-Event-ID")
|
||
try:
|
||
async for entry in bridge.subscribe(record.run_id, last_event_id=last_event_id):
|
||
if record.on_disconnect == DisconnectMode.cancel and await request.is_disconnected():
|
||
break
|
||
|
||
if entry is HEARTBEAT_SENTINEL:
|
||
yield ": heartbeat\n\n"
|
||
continue
|
||
|
||
if entry is END_SENTINEL:
|
||
yield format_sse("end", None, event_id=entry.id or None)
|
||
return
|
||
|
||
yield format_sse(entry.event, entry.data, event_id=entry.id or None)
|
||
|
||
except (ConnectionError, BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
|
||
logger.info("SSE client connection lost run_id=%s thread_id=%s", record.run_id, record.thread_id)
|
||
|
||
finally:
|
||
if record.status in (RunStatus.pending, RunStatus.running):
|
||
if record.on_disconnect == DisconnectMode.cancel:
|
||
await run_mgr.cancel(record.run_id)
|