deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/community/configurable_search/tools.py
2026-09-07 18:24:55 +08:00

354 lines
13 KiB
Python

"""Configurable JSON-over-HTTP implementation of the ``web_search`` tool."""
from __future__ import annotations
import json
import logging
import re
from copy import deepcopy
from dataclasses import dataclass
from datetime import date, datetime, timedelta
from typing import Any
from urllib.parse import quote
from zoneinfo import ZoneInfo
import httpx
from langchain.tools import tool
from deerflow.config import get_app_config
from deerflow.config.tool_config import ToolConfig
logger = logging.getLogger(__name__)
_BEIJING_TZ = ZoneInfo("Asia/Shanghai")
_ISO_DATE_RE = re.compile(r"(?P<year>\d{4})(?:[-/.]|\u5e74)(?P<month>\d{1,2})(?:[-/.]|\u6708)(?P<day>\d{1,2})(?:\u65e5)?")
_RECENT_DAYS_RE = re.compile(r"(\u8fd1\u51e0\u5929|\u6700\u8fd1\u51e0\u5929|\u8fd9\u51e0\u5929|\u8fd1\u4e00\u5468|\u6700\u8fd1\u4e00\u5468|\u8fd17\u5929|\u6700\u8fd17\u5929)")
_RECENT_MONTHS_RE = re.compile(r"(\u8fd1\u671f|\u6700\u8fd1|\u8fd1\u4e09\u4e2a\u6708|\u6700\u8fd1\u4e09\u4e2a\u6708|\u8fd13\u4e2a\u6708|\u6700\u8fd13\u4e2a\u6708)")
_TOP_K_RE = re.compile(r"(?:(?:\u8fd4\u56de|\u67e5|\u68c0\u7d22|\u53d6|\u524d)\s*)?(?P<count>\d{1,3})\s*(?:\u6761|\u6761\u6570|\u7bc7|\u4efd)")
@dataclass(frozen=True)
class ConfigurableSearchSettings:
"""Validated runtime settings read from the ``web_search`` tool entry."""
enabled: bool
endpoint: str
payload: dict[str, Any]
verify_ssl: bool
timeout: float
max_results: int
result_url_template: str
@classmethod
def from_tool_config(cls, config: ToolConfig | None) -> ConfigurableSearchSettings:
if config is None:
raise ValueError("web_search is not configured")
extra = config.model_extra or {}
endpoint = str(extra.get("endpoint") or "").strip()
if not endpoint.startswith(("http://", "https://")):
raise ValueError("web_search.endpoint must use http:// or https://")
raw_payload = extra.get("payload", {})
if not isinstance(raw_payload, dict):
raise ValueError("web_search.payload must be an object")
try:
timeout = float(extra.get("timeout", 30))
except (TypeError, ValueError) as exc:
raise ValueError("web_search.timeout must be a number") from exc
if timeout <= 0:
raise ValueError("web_search.timeout must be greater than zero")
try:
max_results = int(extra.get("max_results", 5))
except (TypeError, ValueError) as exc:
raise ValueError("web_search.max_results must be an integer") from exc
if max_results <= 0:
raise ValueError("web_search.max_results must be greater than zero")
result_url_template = str(extra.get("result_url_template") or "").strip()
if result_url_template and "{recUuid}" not in result_url_template:
raise ValueError("web_search.result_url_template must contain {recUuid}")
verify_ssl = extra.get("verify_ssl", True)
if not isinstance(verify_ssl, bool):
raise ValueError("web_search.verify_ssl must be true or false")
return cls(
enabled=config.enabled,
endpoint=endpoint,
payload=deepcopy(raw_payload),
verify_ssl=verify_ssl,
timeout=timeout,
max_results=max_results,
result_url_template=result_url_template,
)
def _text(value: Any) -> str:
"""Normalize optional scalar result fields without leaking ``None``."""
if value is None:
return ""
if isinstance(value, str):
return value
return str(value)
def _metadata(item: dict[str, Any]) -> dict[str, Any]:
value = item.get("metadata")
return value if isinstance(value, dict) else {}
def _first_text(*values: Any) -> str:
for value in values:
text = _text(value).strip()
if text:
return text
return ""
def _looks_like_url(value: str) -> bool:
return value.startswith(("http://", "https://"))
def _today() -> date:
return datetime.now(_BEIJING_TZ).date()
def _format_date(value: date) -> str:
return value.strftime("%Y-%m-%d")
def _subtract_months(value: date, months: int) -> date:
month_index = value.month - 1 - months
year = value.year + month_index // 12
month = month_index % 12 + 1
leap_year = year % 4 == 0 and (year % 100 != 0 or year % 400 == 0)
month_lengths = [31, 29 if leap_year else 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
return value.replace(year=year, month=month, day=min(value.day, month_lengths[month - 1]))
def _normalize_date(value: str) -> str:
text = str(value or "").strip()
if not text:
return ""
match = _ISO_DATE_RE.search(text)
if not match:
return ""
try:
return _format_date(
date(
int(match.group("year")),
int(match.group("month")),
int(match.group("day")),
)
)
except ValueError:
return ""
def _extract_query_dates(query: str) -> tuple[str, str]:
dates = [_normalize_date(match.group(0)) for match in _ISO_DATE_RE.finditer(query)]
dates = [item for item in dates if item]
if len(dates) >= 2:
return min(dates[0], dates[1]), max(dates[0], dates[1])
if len(dates) == 1:
return dates[0], dates[0]
return "", ""
def _resolve_date_range(
query: str,
*,
start_date: str = "",
end_date: str = "",
today: date | None = None,
) -> tuple[str, str]:
normalized_start = _normalize_date(start_date)
normalized_end = _normalize_date(end_date)
if normalized_start or normalized_end:
return normalized_start, normalized_end
query_start, query_end = _extract_query_dates(query)
if query_start or query_end:
return query_start, query_end
current = today or _today()
if _RECENT_DAYS_RE.search(query):
return _format_date(current - timedelta(days=7)), _format_date(current)
if _RECENT_MONTHS_RE.search(query):
return _format_date(_subtract_months(current, 3)), _format_date(current)
return "", ""
def _resolve_top_k(query: str, top_k: str | int = "") -> str:
explicit = str(top_k).strip() if top_k is not None else ""
if explicit:
return explicit
match = _TOP_K_RE.search(query)
if not match:
return ""
return match.group("count")
def _error_result(query: str, message: str) -> str:
return json.dumps(
{
"query": query,
"total_results": 0,
"results": [],
"error": message,
},
ensure_ascii=False,
)
def _normalize_results(data: Any, settings: ConfigurableSearchSettings) -> list[dict[str, str]]:
if isinstance(data, list):
raw_results = data
elif isinstance(data, dict):
raw_results = data.get("results")
else:
raise ValueError("search API response must be a JSON object or array")
if not isinstance(raw_results, list):
raise ValueError("search API response field 'results' must be an array")
normalized: list[dict[str, str]] = []
for item in raw_results[: settings.max_results]:
if not isinstance(item, dict):
logger.warning("Ignoring non-object search result: %r", item)
continue
metadata = _metadata(item)
rec_uuid = _first_text(item.get("recUuid"), item.get("rec_uuid"), metadata.get("recUuid"), metadata.get("rec_uuid"))
source = _first_text(item.get("source"), item.get("source1"), metadata.get("source"), metadata.get("source1"))
original_url = _first_text(
item.get("url"),
item.get("link"),
item.get("href"),
item.get("source_url"),
metadata.get("url"),
metadata.get("link"),
metadata.get("href"),
metadata.get("source_url"),
)
if not original_url and _looks_like_url(source):
original_url = source
result_url = original_url
if rec_uuid and settings.result_url_template:
result_url = settings.result_url_template.replace(
"{recUuid}",
quote(rec_uuid, safe=""),
)
normalized_item = {
"recUuid": rec_uuid,
"content": _first_text(item.get("content"), item.get("page_content"), metadata.get("content")),
"title": _first_text(item.get("title"), item.get("m_title"), metadata.get("m_title")),
"url": result_url,
}
if source:
normalized_item["source"] = source
publish_time = _first_text(item.get("time"), item.get("publish_time"), item.get("m_publish"), metadata.get("m_publish"))
if publish_time:
normalized_item["time"] = publish_time
normalized_item["publish_time"] = publish_time
normalized.append(normalized_item)
return normalized
async def execute_search(
query: str,
settings: ConfigurableSearchSettings,
*,
start_date: str = "",
end_date: str = "",
top_k: str | int = "",
today: date | None = None,
transport: httpx.AsyncBaseTransport | None = None,
) -> str:
"""Execute a search request; ``transport`` exists for network-free tests."""
normalized_query = query.strip()
if not normalized_query:
return _error_result(query, "Search query must not be empty")
if not settings.enabled:
return _error_result(normalized_query, "web_search is disabled")
payload = deepcopy(settings.payload)
resolved_start, resolved_end = _resolve_date_range(
normalized_query,
start_date=start_date,
end_date=end_date,
today=today,
)
payload["query"] = normalized_query
payload["start_date"] = resolved_start
payload["end_date"] = resolved_end
payload["top_k"] = _resolve_top_k(normalized_query, top_k)
try:
async with httpx.AsyncClient(
timeout=settings.timeout,
verify=settings.verify_ssl,
transport=transport,
) as client:
response = await client.post(settings.endpoint, json=payload)
response.raise_for_status()
data = response.json()
results = _normalize_results(data, settings)
except httpx.TimeoutException:
logger.warning("Configured web search timed out for endpoint %s", settings.endpoint)
return _error_result(normalized_query, "Search API request timed out")
except httpx.HTTPStatusError as exc:
logger.warning(
"Configured web search returned HTTP %s for endpoint %s",
exc.response.status_code,
settings.endpoint,
)
return _error_result(normalized_query, f"Search API returned HTTP {exc.response.status_code}")
except (httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning("Configured web search request failed: %s", exc)
return _error_result(normalized_query, f"Search API request failed: {type(exc).__name__}")
except ValueError as exc:
logger.warning("Configured web search returned an invalid response: %s", exc)
return _error_result(normalized_query, str(exc))
return json.dumps(
{
"query": normalized_query,
"total_results": len(results),
"results": results,
},
ensure_ascii=False,
)
@tool("web_search", parse_docstring=True)
async def web_search_tool(query: str, start_date: str = "", end_date: str = "", top_k: str = "") -> str:
"""Search the internal knowledge base and return source passages.
Args:
query: Knowledge-base retrieval query. Keep the user's subject and time
requirement in this text. If the user says recent/近期/最近, this tool
will search the latest three months. If the user says near days/近几天/
最近几天/这几天/近一周, this tool will search the latest seven days.
start_date: Optional start date in YYYY-MM-DD. Leave empty by default.
Pass this only when the user explicitly asks for a concrete date or
date range.
end_date: Optional end date in YYYY-MM-DD. Leave empty by default.
Pass this only when the user explicitly asks for a concrete date or
date range.
top_k: Optional maximum number of records to return, such as "15".
Leave empty by default. Pass this when the user asks for a specific
number of results.
"""
try:
settings = ConfigurableSearchSettings.from_tool_config(get_app_config().get_tool_config("web_search"))
except ValueError as exc:
return _error_result(query, str(exc))
return await execute_search(query, settings, start_date=start_date, end_date=end_date, top_k=top_k)