354 lines
13 KiB
Python
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)
|