1004 lines
36 KiB
Python
1004 lines
36 KiB
Python
"""Build stable per-turn citation batches from LangGraph messages."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import re
|
||
from typing import Any
|
||
from urllib.parse import quote
|
||
|
||
from deerflow.config.extensions_config import (
|
||
SkillDisplayConfig,
|
||
SkillDisplayResultSetConfig,
|
||
get_extensions_config,
|
||
)
|
||
|
||
_SKILL_PATH_RE = re.compile(r"[\\/]skills[\\/](?:public|custom)[\\/]([^\\/]+)(?:[\\/]|$)")
|
||
|
||
|
||
def build_reference_batches(messages: list[Any]) -> list[dict[str, Any]]:
|
||
batches: list[dict[str, Any]] = []
|
||
for index, message in enumerate(messages):
|
||
if _message_type(message) != "human":
|
||
continue
|
||
end = len(messages)
|
||
for cursor in range(index + 1, len(messages)):
|
||
if _message_type(messages[cursor]) == "human":
|
||
end = cursor
|
||
break
|
||
preloaded_sources = _collect_preloaded_sources(message, start_index=1)
|
||
sources = [
|
||
*preloaded_sources,
|
||
*_collect_sources(messages[index + 1 : end], start_index=len(preloaded_sources) + 1),
|
||
]
|
||
if sources:
|
||
anchor = _find_reference_batch_anchor(messages, index + 1, end)
|
||
message_id = _message_id(anchor) or _message_id(message) or str(index)
|
||
batches.append(
|
||
{
|
||
"id": message_id,
|
||
"assistant_message_id": message_id,
|
||
"sources": sources,
|
||
}
|
||
)
|
||
return batches
|
||
|
||
|
||
def build_latest_reference_batch(messages: list[Any]) -> dict[str, Any] | None:
|
||
"""Return the single combined citation batch for the latest human turn."""
|
||
human_index = _find_latest_human_index(messages)
|
||
if human_index is None:
|
||
return None
|
||
preloaded_sources = _collect_preloaded_sources(messages[human_index], start_index=1)
|
||
sources = [
|
||
*preloaded_sources,
|
||
*_collect_sources(messages[human_index + 1 :], start_index=len(preloaded_sources) + 1),
|
||
]
|
||
if not sources:
|
||
return None
|
||
assistant = _find_reference_batch_anchor(messages, human_index + 1, len(messages))
|
||
message_id = _message_id(assistant) or _message_id(messages[human_index]) or str(human_index)
|
||
return {
|
||
"id": message_id,
|
||
"assistant_message_id": message_id,
|
||
"source_count": len(sources),
|
||
"sources": sources,
|
||
}
|
||
|
||
|
||
def build_latest_reference_prompt(messages: list[Any]) -> tuple[str, dict[str, Any] | None]:
|
||
batch = build_latest_reference_batch(messages)
|
||
if batch is None:
|
||
return "", None
|
||
return build_reference_prompt_for_batch(batch), batch
|
||
|
||
|
||
def build_reference_prompt_for_batch(batch: dict[str, Any]) -> str:
|
||
"""Build the authoritative model prompt from one already-combined batch."""
|
||
sources = [
|
||
source
|
||
for source in batch.get("sources", [])
|
||
if isinstance(source, dict)
|
||
]
|
||
source_count = len(sources)
|
||
lines = [
|
||
f"本轮所有检索结果已合并为同一个 referenceBatch,共 {source_count} 条,编号从 [1] 到 [{source_count}] 唯一连续。",
|
||
"下面这一批就是当前回答唯一可用的参考文献批次;不要按检索词、工具调用次数或工具原始编号重新分批编号。",
|
||
"如果你仍需要补充检索,请只继续调用工具,不要输出正文;最终回答只能引用补充检索后系统重新合并出的最新 referenceBatch。",
|
||
"",
|
||
f"<reference_batch source_count=\"{source_count}\">",
|
||
]
|
||
for source in sources:
|
||
title = source.get("title") or f"参考文献 {source.get('index')}"
|
||
meta = " · ".join(
|
||
str(item)
|
||
for item in (source.get("source"), source.get("author"), source.get("time"))
|
||
if item
|
||
)
|
||
content = str(source.get("content") or source.get("snippet") or "").strip()
|
||
if len(content) > 520:
|
||
content = content[:517] + "..."
|
||
head = f"[{source.get('index')}] {title}"
|
||
if meta:
|
||
head += f"({meta})"
|
||
lines.append(head)
|
||
if content:
|
||
lines.append(content)
|
||
lines.append("")
|
||
lines.append("</reference_batch>")
|
||
lines.append("")
|
||
lines.extend(
|
||
[
|
||
"引用规则:",
|
||
"1. 回答正文中的关键结论必须标注对应编号,例如 [1]、[2]。",
|
||
"2. 多个来源可写为 [1][3][5]。",
|
||
"3. 只能引用上方存在的编号,不要编造编号。",
|
||
"4. 不要使用工具原始编号,也不要重新排序、重新分批或重新从 [1] 开始编号。",
|
||
"5. 不要在正文末尾堆完整参考文献列表,前端会展示右侧参考文献。",
|
||
"6. 参考文献模式下不要输出 <skill-display .../> 占位符,也不要在正文里额外渲染卡片、表格或列表;前端会在右侧参考文献区域展示。",
|
||
]
|
||
)
|
||
return "\n".join(lines).strip()
|
||
|
||
|
||
def _find_latest_human_index(messages: list[Any]) -> int | None:
|
||
for index in range(len(messages) - 1, -1, -1):
|
||
if _message_type(messages[index]) == "human":
|
||
return index
|
||
return None
|
||
|
||
|
||
def build_reference_debug_summary(messages: list[Any], batches: list[dict[str, Any]] | None = None) -> dict[str, Any]:
|
||
"""Return compact counts for citation extraction logs."""
|
||
turn_summaries: list[dict[str, Any]] = []
|
||
total_tool_messages = 0
|
||
total_raw_items = 0
|
||
total_normalized_items = 0
|
||
total_kept_items = 0
|
||
for index, message in enumerate(messages):
|
||
if _message_type(message) != "human":
|
||
continue
|
||
end = len(messages)
|
||
for cursor in range(index + 1, len(messages)):
|
||
if _message_type(messages[cursor]) == "human":
|
||
end = cursor
|
||
break
|
||
tool_messages = 0
|
||
raw_items = 0
|
||
normalized_items = 0
|
||
kept_items = 0
|
||
for tool_message in messages[index + 1 : end]:
|
||
if _message_type(tool_message) != "tool":
|
||
continue
|
||
tool_messages += 1
|
||
items = _first_list(_extract_tool_payload(tool_message))
|
||
raw_items += len(items)
|
||
for raw_item in items:
|
||
normalized = _normalize_reference_item(raw_item, normalized_items + 1)
|
||
if not normalized:
|
||
continue
|
||
normalized_items += 1
|
||
kept_items += 1
|
||
total_tool_messages += tool_messages
|
||
total_raw_items += raw_items
|
||
total_normalized_items += normalized_items
|
||
total_kept_items += kept_items
|
||
turn_summaries.append(
|
||
{
|
||
"human_index": index,
|
||
"tool_messages": tool_messages,
|
||
"raw_items": raw_items,
|
||
"normalized_items": normalized_items,
|
||
"kept_items": kept_items,
|
||
}
|
||
)
|
||
batch_summaries = [
|
||
{
|
||
"id": batch.get("id"),
|
||
"assistant_message_id": batch.get("assistant_message_id"),
|
||
"source_count": len(batch.get("sources") or []),
|
||
}
|
||
for batch in (batches or build_reference_batches(messages))
|
||
]
|
||
return {
|
||
"messages": len(messages),
|
||
"turns": len(turn_summaries),
|
||
"tool_messages": total_tool_messages,
|
||
"raw_items": total_raw_items,
|
||
"normalized_items": total_normalized_items,
|
||
"kept_items": total_kept_items,
|
||
"turns_detail": turn_summaries,
|
||
"batches": batch_summaries,
|
||
}
|
||
|
||
|
||
def build_reference_debug_payload(
|
||
*,
|
||
messages: list[Any] | None = None,
|
||
batch: dict[str, Any] | None = None,
|
||
batches: list[dict[str, Any]] | None = None,
|
||
prompt: str | None = None,
|
||
) -> dict[str, Any]:
|
||
"""Return a JSON-safe payload for comparing model prompts and UI batches."""
|
||
payload: dict[str, Any] = {}
|
||
if messages is not None:
|
||
payload["message_trace"] = [_message_debug_item(index, message) for index, message in enumerate(messages)]
|
||
if prompt is not None:
|
||
payload["model_reference_prompt"] = prompt
|
||
if batch is not None:
|
||
payload["batch"] = _batch_debug_item(batch)
|
||
if batches is not None:
|
||
payload["batches"] = [_batch_debug_item(item) for item in batches]
|
||
return payload
|
||
|
||
|
||
def adapt_reference_batches_with_displays(
|
||
batches: list[dict[str, Any]],
|
||
displays: dict[str, SkillDisplayConfig | dict[str, Any]],
|
||
) -> list[dict[str, Any]]:
|
||
"""Apply persisted skill display configs to already-built reference batches.
|
||
|
||
The reference builder can run before the gateway has a fresh extensions
|
||
cache, while the skill adapter editor persists the user's latest mapping in
|
||
the database. This pass lets the API response use that DB-backed mapping
|
||
without changing citation indexes.
|
||
"""
|
||
normalized_displays: dict[str, SkillDisplayConfig] = {}
|
||
for skill_name, display in displays.items():
|
||
try:
|
||
config = (
|
||
display
|
||
if isinstance(display, SkillDisplayConfig)
|
||
else SkillDisplayConfig.model_validate(display)
|
||
)
|
||
except Exception:
|
||
continue
|
||
has_enabled_result_set = any(
|
||
result_set.enabled and result_set.mode != "none"
|
||
for result_set in config.result_sets
|
||
)
|
||
if not config.enabled or (config.mode == "none" and not has_enabled_result_set):
|
||
continue
|
||
normalized_displays[skill_name] = config
|
||
|
||
if not normalized_displays:
|
||
return batches
|
||
|
||
displays_by_normalized_name = {
|
||
_normalize_name(name): (name, display)
|
||
for name, display in normalized_displays.items()
|
||
}
|
||
adapted_batches: list[dict[str, Any]] = []
|
||
for batch in batches:
|
||
sources: list[dict[str, Any]] = []
|
||
for source in batch.get("sources", []):
|
||
if not isinstance(source, dict):
|
||
continue
|
||
skill_name = str(source.get("skillName") or "")
|
||
display = normalized_displays.get(skill_name)
|
||
if display is None:
|
||
normalized_source = _normalize_name(skill_name)
|
||
normalized_match = displays_by_normalized_name.get(normalized_source)
|
||
if normalized_match is None:
|
||
for normalized_skill, candidate in displays_by_normalized_name.items():
|
||
if (
|
||
normalized_source.endswith(normalized_skill)
|
||
or normalized_skill in normalized_source
|
||
):
|
||
normalized_match = candidate
|
||
break
|
||
if normalized_match is not None:
|
||
skill_name, display = normalized_match
|
||
raw = source.get("raw")
|
||
if display is not None and isinstance(raw, dict):
|
||
index_value = int(source.get("index") or len(sources) + 1)
|
||
display_for_source: SkillDisplayConfig | SkillDisplayResultSetConfig = display
|
||
result_set_key = source.get("resultSetKey")
|
||
if result_set_key:
|
||
display_for_source = next(
|
||
(
|
||
result_set
|
||
for result_set in display.result_sets
|
||
if result_set.key == result_set_key
|
||
),
|
||
display,
|
||
)
|
||
elif display.result_sets and display.mode == "none":
|
||
sources.append(source)
|
||
continue
|
||
configured = _normalize_with_display_config(
|
||
raw,
|
||
index_value,
|
||
skill_name,
|
||
display_for_source,
|
||
)
|
||
if configured is not None:
|
||
if isinstance(display_for_source, SkillDisplayResultSetConfig):
|
||
configured["resultSetKey"] = display_for_source.key
|
||
configured["resultSetLabel"] = display_for_source.label
|
||
sources.append(
|
||
{
|
||
**source,
|
||
**configured,
|
||
"index": index_value,
|
||
"skillName": skill_name,
|
||
"raw": raw,
|
||
}
|
||
)
|
||
continue
|
||
sources.append(source)
|
||
adapted_batches.append({**batch, "sources": sources})
|
||
return adapted_batches
|
||
|
||
|
||
def _message_debug_item(index: int, message: Any) -> dict[str, Any]:
|
||
content = _message_content(message)
|
||
tool_calls = _extract_tool_calls(message)
|
||
return {
|
||
"position": index,
|
||
"type": _message_type(message),
|
||
"id": _message_id(message),
|
||
"name": _message_name(message) if _message_type(message) == "tool" else None,
|
||
"tool_call_id": _message_tool_call_id(message) if _message_type(message) == "tool" else None,
|
||
"tool_calls": [
|
||
{
|
||
"id": call.get("id"),
|
||
"name": call.get("name"),
|
||
"args": call.get("args"),
|
||
}
|
||
for call in tool_calls
|
||
],
|
||
"content_length": len(content) if isinstance(content, str) else None,
|
||
}
|
||
|
||
|
||
def _batch_debug_item(batch: dict[str, Any]) -> dict[str, Any]:
|
||
sources = [source for source in batch.get("sources", []) if isinstance(source, dict)]
|
||
return {
|
||
"id": batch.get("id"),
|
||
"assistant_message_id": batch.get("assistant_message_id"),
|
||
"source_count": len(sources),
|
||
"sources": [_source_debug_item(source) for source in sources],
|
||
}
|
||
|
||
|
||
def _source_debug_item(source: dict[str, Any]) -> dict[str, Any]:
|
||
raw = source.get("raw")
|
||
raw_record = raw if isinstance(raw, dict) else {}
|
||
return {
|
||
"index": source.get("index"),
|
||
"id": source.get("id"),
|
||
"skillName": source.get("skillName"),
|
||
"mode": source.get("mode"),
|
||
"resultSetKey": source.get("resultSetKey"),
|
||
"resultSetLabel": source.get("resultSetLabel"),
|
||
"title": source.get("title"),
|
||
"source": source.get("source"),
|
||
"author": source.get("author"),
|
||
"time": source.get("time"),
|
||
"url": source.get("url"),
|
||
"snippet": source.get("snippet"),
|
||
"content": source.get("content"),
|
||
"raw_id_candidates": {
|
||
"recUuid": raw_record.get("recUuid"),
|
||
"rec_uuid": raw_record.get("rec_uuid"),
|
||
"uuid": raw_record.get("uuid"),
|
||
"docId": raw_record.get("docId"),
|
||
"doc_id": raw_record.get("doc_id"),
|
||
"id": raw_record.get("id"),
|
||
"url": raw_record.get("url"),
|
||
},
|
||
"raw": raw,
|
||
}
|
||
|
||
|
||
def _collect_sources(messages: list[Any], *, start_index: int = 1) -> list[dict[str, Any]]:
|
||
sources: list[dict[str, Any]] = []
|
||
pending_tool_contexts: list[dict[str, Any]] = []
|
||
tool_contexts_by_id: dict[str, dict[str, Any]] = {}
|
||
for message in messages:
|
||
message_type = _message_type(message)
|
||
if message_type == "ai":
|
||
for tool_call in _extract_tool_calls(message):
|
||
context = _build_tool_context(tool_call)
|
||
pending_tool_contexts.append(context)
|
||
if context.get("id"):
|
||
tool_contexts_by_id[str(context["id"])] = context
|
||
continue
|
||
if message_type != "tool":
|
||
continue
|
||
tool_context = _resolve_tool_context(
|
||
message,
|
||
pending_tool_contexts,
|
||
tool_contexts_by_id,
|
||
)
|
||
payload = _extract_tool_payload(message)
|
||
configured_sources = _collect_configured_sources(
|
||
payload,
|
||
len(sources) + 1,
|
||
tool_context,
|
||
)
|
||
if configured_sources is not None:
|
||
for normalized in configured_sources:
|
||
normalized["index"] = start_index + len(sources)
|
||
sources.append(normalized)
|
||
continue
|
||
for raw_item in _first_list(payload):
|
||
normalized = _normalize_reference_item(
|
||
raw_item,
|
||
len(sources) + 1,
|
||
payload=payload,
|
||
tool_context=tool_context,
|
||
)
|
||
if not normalized:
|
||
continue
|
||
normalized["index"] = start_index + len(sources)
|
||
sources.append(normalized)
|
||
return sources
|
||
|
||
|
||
def _collect_preloaded_sources(message: Any, *, start_index: int = 1) -> list[dict[str, Any]]:
|
||
extra = _message_additional_kwargs(message)
|
||
raw_sources = extra.get("llmwiki_reference_sources")
|
||
if not isinstance(raw_sources, list):
|
||
return []
|
||
sources: list[dict[str, Any]] = []
|
||
for raw in raw_sources:
|
||
if not isinstance(raw, dict):
|
||
continue
|
||
source = dict(raw)
|
||
index = start_index + len(sources)
|
||
source["index"] = index
|
||
source.setdefault("skillName", "llmwiki_search")
|
||
source.setdefault("mode", "citation")
|
||
source.setdefault("type", "llmwiki")
|
||
title = str(source.get("title") or "").strip()
|
||
if not title:
|
||
source["title"] = f"LLMWiki 片段 {index}"
|
||
snippet = str(source.get("snippet") or source.get("content") or "").strip()
|
||
source["snippet"] = snippet[:360]
|
||
if "raw" not in source or not isinstance(source.get("raw"), dict):
|
||
source["raw"] = {key: value for key, value in source.items() if key != "raw"}
|
||
sources.append(source)
|
||
return sources
|
||
|
||
|
||
def _collect_configured_sources(
|
||
payload: Any,
|
||
fallback_index: int,
|
||
tool_context: dict[str, Any] | None,
|
||
) -> list[dict[str, Any]] | None:
|
||
display_match = _resolve_display_match({}, payload, tool_context)
|
||
if display_match is None:
|
||
return None
|
||
skill_name, display = display_match
|
||
configs: list[SkillDisplayConfig | SkillDisplayResultSetConfig] = [
|
||
result_set
|
||
for result_set in display.result_sets
|
||
if result_set.enabled and result_set.mode != "none"
|
||
]
|
||
if not configs:
|
||
configs = [display]
|
||
|
||
sources: list[dict[str, Any]] = []
|
||
for config in configs:
|
||
raw_items = _list_at_path(payload, config.result_path)
|
||
if raw_items is None:
|
||
continue
|
||
for raw_item in raw_items:
|
||
normalized = _normalize_with_display_config(
|
||
raw_item,
|
||
fallback_index + len(sources),
|
||
skill_name,
|
||
config,
|
||
)
|
||
if normalized is None:
|
||
continue
|
||
if isinstance(config, SkillDisplayResultSetConfig):
|
||
normalized["resultSetKey"] = config.key
|
||
normalized["resultSetLabel"] = config.label
|
||
sources.append(normalized)
|
||
return sources
|
||
|
||
|
||
def _find_reference_batch_anchor(messages: list[Any], start: int, end: int) -> Any | None:
|
||
answer = _find_last_assistant_answer(messages, start, end)
|
||
if answer is not None:
|
||
return answer
|
||
for index in range(end - 1, start - 1, -1):
|
||
candidate = messages[index]
|
||
if _message_type(candidate) == "ai" and _message_id(candidate):
|
||
return candidate
|
||
return None
|
||
|
||
|
||
def _find_last_assistant_answer(messages: list[Any], start: int, end: int) -> Any | None:
|
||
for index in range(end - 1, start - 1, -1):
|
||
candidate = messages[index]
|
||
if (
|
||
_message_type(candidate) == "ai"
|
||
and _message_id(candidate)
|
||
and _has_answer_content(candidate)
|
||
):
|
||
return candidate
|
||
return None
|
||
|
||
|
||
def _message_type(message: Any) -> str:
|
||
if isinstance(message, dict):
|
||
return str(message.get("type") or message.get("role") or "")
|
||
return str(getattr(message, "type", "") or "")
|
||
|
||
|
||
def _message_id(message: Any) -> str | None:
|
||
value = message.get("id") if isinstance(message, dict) else getattr(message, "id", None)
|
||
return str(value) if value else None
|
||
|
||
|
||
def _message_content(message: Any) -> Any:
|
||
return message.get("content") if isinstance(message, dict) else getattr(message, "content", None)
|
||
|
||
|
||
def _message_additional_kwargs(message: Any) -> dict[str, Any]:
|
||
if isinstance(message, dict):
|
||
value = message.get("additional_kwargs") or message.get("additionalKwargs")
|
||
else:
|
||
value = getattr(message, "additional_kwargs", None)
|
||
return value if isinstance(value, dict) else {}
|
||
|
||
|
||
def _message_name(message: Any) -> str:
|
||
if isinstance(message, dict):
|
||
return str(message.get("name") or (message.get("additional_kwargs") or {}).get("name") or "tool")
|
||
return str(getattr(message, "name", None) or "tool")
|
||
|
||
|
||
def _message_tool_call_id(message: Any) -> str | None:
|
||
if isinstance(message, dict):
|
||
value = message.get("tool_call_id") or message.get("toolCallId")
|
||
else:
|
||
value = getattr(message, "tool_call_id", None) or getattr(message, "toolCallId", None)
|
||
return str(value) if value else None
|
||
|
||
|
||
def _has_tool_calls(message: Any) -> bool:
|
||
if isinstance(message, dict):
|
||
direct = message.get("tool_calls")
|
||
chunks = message.get("tool_call_chunks")
|
||
extra = (message.get("additional_kwargs") or {}).get("tool_calls")
|
||
else:
|
||
direct = getattr(message, "tool_calls", None)
|
||
chunks = getattr(message, "tool_call_chunks", None)
|
||
extra = (getattr(message, "additional_kwargs", None) or {}).get("tool_calls")
|
||
return bool(direct or chunks or extra)
|
||
|
||
|
||
def _has_answer_content(message: Any) -> bool:
|
||
content = _message_content(message)
|
||
if isinstance(content, str):
|
||
return bool(content.strip())
|
||
if isinstance(content, list):
|
||
return any(
|
||
isinstance(part, dict)
|
||
and part.get("type") == "text"
|
||
and str(part.get("text") or "").strip()
|
||
for part in content
|
||
)
|
||
return False
|
||
|
||
|
||
def _extract_tool_calls(message: Any) -> list[dict[str, Any]]:
|
||
candidates: list[Any] = []
|
||
if isinstance(message, dict):
|
||
for key in ("tool_calls", "tool_call_chunks"):
|
||
value = message.get(key)
|
||
if isinstance(value, list):
|
||
candidates.extend(value)
|
||
extra = (message.get("additional_kwargs") or {}).get("tool_calls")
|
||
if isinstance(extra, list):
|
||
candidates.extend(extra)
|
||
else:
|
||
for key in ("tool_calls", "tool_call_chunks"):
|
||
value = getattr(message, key, None)
|
||
if isinstance(value, list):
|
||
candidates.extend(value)
|
||
extra = (getattr(message, "additional_kwargs", None) or {}).get("tool_calls")
|
||
if isinstance(extra, list):
|
||
candidates.extend(extra)
|
||
|
||
calls: list[dict[str, Any]] = []
|
||
seen_ids: set[str] = set()
|
||
for raw in candidates:
|
||
call = _normalize_tool_call(raw)
|
||
if not call:
|
||
continue
|
||
call_id = str(call.get("id") or "")
|
||
if call_id and call_id in seen_ids:
|
||
continue
|
||
if call_id:
|
||
seen_ids.add(call_id)
|
||
calls.append(call)
|
||
return calls
|
||
|
||
|
||
def _normalize_tool_call(raw: Any) -> dict[str, Any] | None:
|
||
if not isinstance(raw, dict):
|
||
return None
|
||
fn = raw.get("function") if isinstance(raw.get("function"), dict) else {}
|
||
name = raw.get("name") or fn.get("name") or raw.get("tool_name")
|
||
if not name:
|
||
return None
|
||
return {
|
||
"id": raw.get("id") or raw.get("tool_call_id") or raw.get("call_id"),
|
||
"name": str(name),
|
||
"args": _normalize_tool_args(raw.get("args") or raw.get("arguments") or fn.get("arguments")),
|
||
}
|
||
|
||
|
||
def _normalize_tool_args(value: Any) -> dict[str, Any]:
|
||
if isinstance(value, dict):
|
||
return value
|
||
if isinstance(value, str):
|
||
try:
|
||
parsed = json.loads(value)
|
||
if isinstance(parsed, dict):
|
||
return parsed
|
||
except Exception:
|
||
pass
|
||
return {"command": value}
|
||
return {}
|
||
|
||
|
||
def _build_tool_context(tool_call: dict[str, Any]) -> dict[str, Any]:
|
||
tool_name = str(tool_call.get("name") or "tool")
|
||
args = tool_call.get("args") if isinstance(tool_call.get("args"), dict) else {}
|
||
return {
|
||
"id": tool_call.get("id"),
|
||
"tool_name": tool_name,
|
||
"args": args,
|
||
"skill_name": _infer_skill_name(tool_name, args),
|
||
}
|
||
|
||
|
||
def _resolve_tool_context(
|
||
message: Any,
|
||
pending_tool_contexts: list[dict[str, Any]],
|
||
tool_contexts_by_id: dict[str, dict[str, Any]],
|
||
) -> dict[str, Any]:
|
||
tool_call_id = _message_tool_call_id(message)
|
||
if tool_call_id and tool_call_id in tool_contexts_by_id:
|
||
context = tool_contexts_by_id[tool_call_id]
|
||
pending_tool_contexts[:] = [
|
||
item for item in pending_tool_contexts if str(item.get("id") or "") != tool_call_id
|
||
]
|
||
return context
|
||
if pending_tool_contexts:
|
||
return pending_tool_contexts.pop(0)
|
||
tool_name = _message_name(message)
|
||
return {
|
||
"id": tool_call_id,
|
||
"tool_name": tool_name,
|
||
"args": {},
|
||
"skill_name": _infer_skill_name(tool_name, {}),
|
||
}
|
||
|
||
|
||
def _extract_tool_payload(message: Any) -> Any:
|
||
content = _message_content(message)
|
||
if isinstance(content, dict):
|
||
return _unwrap_tool_payload(content)
|
||
if isinstance(content, list):
|
||
texts: list[str] = []
|
||
for part in content:
|
||
if not isinstance(part, dict):
|
||
continue
|
||
if isinstance(part.get("json"), dict):
|
||
return _unwrap_tool_payload(part["json"])
|
||
if part.get("type") == "json" and isinstance(part.get("json"), dict):
|
||
return _unwrap_tool_payload(part["json"])
|
||
if part.get("type") == "text" and isinstance(part.get("text"), str):
|
||
texts.append(part["text"])
|
||
unwrapped = _unwrap_tool_payload(part)
|
||
if unwrapped is not None:
|
||
return unwrapped
|
||
content = "\n".join(texts)
|
||
if isinstance(content, str):
|
||
return _parse_json_like(content)
|
||
return None
|
||
|
||
|
||
def _unwrap_tool_payload(value: dict[str, Any]) -> Any:
|
||
list_keys = ("results", "data", "items", "hits", "documents", "records", "sources")
|
||
if any(isinstance(value.get(key), list) for key in list_keys):
|
||
return value
|
||
for key in ("result", "output", "stdout", "content", "text", "payload", "response"):
|
||
nested = value.get(key)
|
||
if isinstance(nested, str):
|
||
parsed = _parse_json_like(nested)
|
||
if parsed is not None:
|
||
return parsed
|
||
if isinstance(nested, dict):
|
||
unwrapped = _unwrap_tool_payload(nested)
|
||
if unwrapped is not None:
|
||
return unwrapped
|
||
return value
|
||
|
||
|
||
def _parse_json_like(text: str) -> Any:
|
||
trimmed = text.strip()
|
||
if not trimmed:
|
||
return None
|
||
try:
|
||
return json.loads(trimmed)
|
||
except Exception:
|
||
pass
|
||
first = trimmed.find("{")
|
||
last = trimmed.rfind("}")
|
||
if 0 <= first < last:
|
||
try:
|
||
return json.loads(trimmed[first : last + 1])
|
||
except Exception:
|
||
return None
|
||
return None
|
||
|
||
|
||
def _first_list(payload: Any) -> list[Any]:
|
||
if isinstance(payload, list):
|
||
return payload
|
||
if not isinstance(payload, dict):
|
||
return []
|
||
for key in ("results", "data", "items", "hits", "documents", "records", "entries", "matches", "chunks", "sources"):
|
||
value = payload.get(key)
|
||
if isinstance(value, list):
|
||
return value
|
||
for key in ("result", "payload", "response", "output", "data"):
|
||
found = _first_list(payload.get(key))
|
||
if found:
|
||
return found
|
||
return []
|
||
|
||
|
||
def _list_at_path(payload: Any, path: str | None) -> list[dict[str, Any]] | None:
|
||
cursor = payload
|
||
if path:
|
||
for part in path.split("."):
|
||
if isinstance(cursor, dict):
|
||
cursor = cursor.get(part)
|
||
else:
|
||
return None
|
||
if not isinstance(cursor, list):
|
||
return None
|
||
return [item for item in cursor if isinstance(item, dict)]
|
||
|
||
|
||
def _nested_records(raw: Any) -> list[dict[str, Any]]:
|
||
if not isinstance(raw, dict):
|
||
return []
|
||
layers = [raw]
|
||
for key in ("metadata", "meta", "attrs", "attributes", "properties", "extra"):
|
||
nested = raw.get(key)
|
||
if isinstance(nested, dict):
|
||
layers.append(nested)
|
||
return layers
|
||
|
||
|
||
def _pick(layers: list[dict[str, Any]], keys: tuple[str, ...]) -> str | None:
|
||
for layer in layers:
|
||
for key in keys:
|
||
value = layer.get(key)
|
||
if isinstance(value, (str, int, float)) and str(value).strip():
|
||
return str(value).strip()
|
||
if isinstance(value, list):
|
||
text = _stringify(value, "")
|
||
if text:
|
||
return text
|
||
return None
|
||
|
||
|
||
def _enabled_skill_displays() -> dict[str, SkillDisplayConfig]:
|
||
try:
|
||
return {
|
||
name: state.display
|
||
for name, state in get_extensions_config().skills.items()
|
||
if state.display is not None
|
||
and state.enabled
|
||
and state.display.enabled
|
||
and (
|
||
state.display.mode != "none"
|
||
or any(
|
||
result_set.enabled and result_set.mode != "none"
|
||
for result_set in state.display.result_sets
|
||
)
|
||
)
|
||
}
|
||
except Exception:
|
||
return {}
|
||
|
||
|
||
def _infer_skill_name(tool_name: str, args: dict[str, Any]) -> str | None:
|
||
displays = _enabled_skill_displays()
|
||
if tool_name in displays:
|
||
return tool_name
|
||
normalized_tool = _normalize_name(tool_name)
|
||
for skill_name in displays:
|
||
normalized_skill = _normalize_name(skill_name)
|
||
if (
|
||
normalized_tool == normalized_skill
|
||
or normalized_tool.endswith(normalized_skill)
|
||
or normalized_skill in normalized_tool
|
||
):
|
||
return skill_name
|
||
for value in args.values():
|
||
if not isinstance(value, str):
|
||
continue
|
||
match = _SKILL_PATH_RE.search(value)
|
||
if match:
|
||
return match.group(1)
|
||
return None
|
||
|
||
|
||
def _resolve_display_match(
|
||
raw: dict[str, Any],
|
||
payload: Any,
|
||
tool_context: dict[str, Any] | None,
|
||
) -> tuple[str, SkillDisplayConfig] | None:
|
||
displays = _enabled_skill_displays()
|
||
if not displays:
|
||
return None
|
||
skill_name = str((tool_context or {}).get("skill_name") or "")
|
||
display = displays.get(skill_name)
|
||
if display is not None:
|
||
return skill_name, display
|
||
tool_name = str((tool_context or {}).get("tool_name") or "")
|
||
inferred = _infer_skill_name(tool_name, (tool_context or {}).get("args") or {})
|
||
if inferred and inferred in displays:
|
||
return inferred, displays[inferred]
|
||
return None
|
||
|
||
|
||
def _normalize_with_display_config(
|
||
raw: dict[str, Any],
|
||
fallback_index: int,
|
||
skill_name: str,
|
||
display: SkillDisplayConfig | SkillDisplayResultSetConfig,
|
||
) -> dict[str, Any] | None:
|
||
title = _stringify(_get_by_path(raw, display.title_field), f"结果 {fallback_index}")
|
||
snippet = _stringify(_get_by_path(raw, display.summary_field), "")
|
||
content = _stringify(_get_by_path(raw, display.content_field), snippet)
|
||
source = _stringify(_get_by_path(raw, display.source_field), "") or None
|
||
author = _stringify(_get_by_path(raw, display.author_field), "") or None
|
||
item_type = _stringify(_get_by_path(raw, display.type_field), "result") or "result"
|
||
time = _stringify(_get_by_path(raw, display.time_field), "") or None
|
||
item_id = _stringify(_get_by_path(raw, display.id_field), "") or None
|
||
score = _number_or_none(_get_by_path(raw, display.score_field))
|
||
url = _resolve_display_link(raw, display)
|
||
if not title and not content and not url:
|
||
return None
|
||
return {
|
||
"id": item_id or url,
|
||
"skillName": skill_name,
|
||
"mode": display.mode,
|
||
"title": title,
|
||
"type": item_type,
|
||
"snippet": snippet,
|
||
"content": content,
|
||
"source": source,
|
||
"author": author,
|
||
"time": time,
|
||
"url": url,
|
||
"score": score,
|
||
"tableColumns": [column.model_dump(by_alias=True) for column in display.table_columns],
|
||
"raw": raw,
|
||
}
|
||
|
||
|
||
def _resolve_display_link(
|
||
raw: dict[str, Any],
|
||
display: SkillDisplayConfig | SkillDisplayResultSetConfig,
|
||
) -> str | None:
|
||
url = _stringify(_get_by_path(raw, display.url_field), "")
|
||
item_id = _stringify(_get_by_path(raw, display.id_field), "")
|
||
if display.link_strategy == "url":
|
||
return url or None
|
||
if display.link_strategy == "id":
|
||
return _render_template(display.link_template, raw) if item_id and display.link_template else None
|
||
return url or (_render_template(display.link_template, raw) if item_id and display.link_template else None) or None
|
||
|
||
|
||
def _render_template(template: str | None, raw: dict[str, Any]) -> str | None:
|
||
if not template:
|
||
return None
|
||
|
||
def replace(match: re.Match[str]) -> str:
|
||
value = _get_by_path(raw, match.group(1).strip())
|
||
return quote(_stringify(value, ""), safe="")
|
||
|
||
return re.sub(r"\{([^}]+)\}", replace, template)
|
||
|
||
|
||
def _get_by_path(value: Any, path: str | None) -> Any:
|
||
if not path:
|
||
return None
|
||
cursor = value
|
||
for part in path.split("."):
|
||
if isinstance(cursor, dict):
|
||
cursor = cursor.get(part)
|
||
else:
|
||
return None
|
||
return cursor
|
||
|
||
|
||
def _stringify(value: Any, fallback: str = "") -> str:
|
||
if value is None:
|
||
return fallback
|
||
if isinstance(value, (str, int, float)):
|
||
text = str(value).strip()
|
||
return text if text else fallback
|
||
if isinstance(value, list):
|
||
parts = [
|
||
str(item).strip()
|
||
for item in value
|
||
if isinstance(item, (str, int, float)) and str(item).strip()
|
||
]
|
||
return "、".join(parts) if parts else fallback
|
||
return fallback
|
||
|
||
|
||
def _number_or_none(value: Any) -> float | None:
|
||
if isinstance(value, (int, float)):
|
||
return float(value)
|
||
if isinstance(value, str):
|
||
try:
|
||
return float(value)
|
||
except Exception:
|
||
return None
|
||
return None
|
||
|
||
|
||
def _normalize_name(name: str) -> str:
|
||
return re.sub(r"[^a-z0-9]+", "", name.lower())
|
||
|
||
|
||
def _normalize_reference_item(
|
||
raw: Any,
|
||
fallback_index: int,
|
||
*,
|
||
payload: Any = None,
|
||
tool_context: dict[str, Any] | None = None,
|
||
) -> dict[str, Any] | None:
|
||
if not isinstance(raw, dict):
|
||
return None
|
||
display_match = _resolve_display_match(raw, payload, tool_context)
|
||
if display_match is not None:
|
||
skill_name, display = display_match
|
||
configured = _normalize_with_display_config(raw, fallback_index, skill_name, display)
|
||
if configured is not None:
|
||
return configured
|
||
layers = _nested_records(raw)
|
||
title = _pick(layers, ("title", "m_title", "name", "question", "keyword", "query", "text"))
|
||
content = _pick(layers, ("content", "page_content", "summary", "snippet", "description", "content_preview", "abstract", "body"))
|
||
url = _pick(layers, ("url", "link", "href", "source_url", "web_url"))
|
||
source = _pick(layers, ("source", "source1", "site", "provider", "from"))
|
||
author = _pick(layers, ("author", "authors", "creator", "writer", "byline", "publisher", "owner"))
|
||
time = _pick(layers, ("time", "date", "published_date", "publishedDate", "m_publish", "publish_time", "created_at"))
|
||
item_id = _pick(layers, ("recUuid", "rec_uuid", "uuid", "docId", "doc_id", "id", "recordId", "taskId"))
|
||
score_raw = _pick(layers, ("score", "relevance", "similarity"))
|
||
score = None
|
||
if score_raw is not None:
|
||
try:
|
||
score = float(score_raw)
|
||
except Exception:
|
||
score = None
|
||
if not title:
|
||
title = source or f"结果 {fallback_index}"
|
||
if not content:
|
||
content = ""
|
||
if not title and not content and not url:
|
||
return None
|
||
return {
|
||
"id": item_id or url,
|
||
"skillName": str((tool_context or {}).get("skill_name") or (tool_context or {}).get("tool_name") or "backend-reference"),
|
||
"mode": "citation",
|
||
"title": title,
|
||
"type": "result",
|
||
"snippet": content,
|
||
"content": content,
|
||
"source": source,
|
||
"author": author,
|
||
"time": time,
|
||
"url": url,
|
||
"score": score,
|
||
"raw": raw,
|
||
}
|
||
|
||
|
||
def _reference_key(source: dict[str, Any]) -> str:
|
||
raw = source.get("raw") if isinstance(source.get("raw"), dict) else {}
|
||
item_id = source.get("id") or source.get("url") or raw.get("recUuid") or raw.get("uuid") or raw.get("docId") or raw.get("id")
|
||
if item_id:
|
||
return str(item_id).casefold().strip()
|
||
return "|".join(
|
||
str(source.get(key) or "").casefold().strip()
|
||
for key in ("title", "source", "time", "content", "snippet")
|
||
)
|