feat:更新智能体工具
This commit is contained in:
338
tools/base.py
Normal file
338
tools/base.py
Normal file
@@ -0,0 +1,338 @@
|
||||
from langchain_core.tools import tool
|
||||
import requests
|
||||
import pandas as pd
|
||||
from cache import _GLOBAL_RAW_CACHE, _GLOBAL_FILTERED_CACHE
|
||||
from config.domain import get_domain
|
||||
|
||||
MAX_DETAIL_COUNT = 20
|
||||
MAX_CONFIRM_COUNT = 50
|
||||
|
||||
|
||||
@tool
|
||||
def query_data(domain: str, start_date: str, end_date: str) -> str:
|
||||
"""
|
||||
按时间范围查询任意业务域的原始数据。
|
||||
|
||||
参数:
|
||||
- domain: 业务域名称(如 repair)
|
||||
- start_date: YYYY-MM-DD
|
||||
- end_date: YYYY-MM-DD
|
||||
"""
|
||||
print(f"[TOOL] query_data | domain={domain} | {start_date} ~ {end_date}")
|
||||
|
||||
cfg = get_domain(domain)
|
||||
token = cfg.token_getter()
|
||||
|
||||
print(f"[TOOL] query_data | api_url={cfg.api_url}")
|
||||
|
||||
resp = requests.get(
|
||||
cfg.api_url,
|
||||
params={"startDate": start_date, "endDate": end_date},
|
||||
headers={cfg.auth_header: token},
|
||||
timeout=15,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
_GLOBAL_RAW_CACHE[domain] = data
|
||||
_GLOBAL_FILTERED_CACHE[domain] = data
|
||||
|
||||
print(f"[TOOL] query_data | fetched={len(data)} rows")
|
||||
return f"已获取 {len(data)} 条 {cfg.entity_name} 数据({start_date} ~ {end_date})"
|
||||
|
||||
|
||||
@tool
|
||||
def filter_data(domain: str, filters: list[dict] | None = None, search: list[str] | None = None) -> str:
|
||||
"""
|
||||
对数据进行筛选与检索。
|
||||
|
||||
参数:
|
||||
- domain: 业务域名称
|
||||
- filters: 结构化筛选条件(AND)
|
||||
- search: 关键词检索(OR)
|
||||
|
||||
示例:
|
||||
{
|
||||
"domain": "repair",
|
||||
"filters": [{"field": "deviceModel", "op": "eq", "value": "销售机"}],
|
||||
"search": ["钢丝绳"]
|
||||
}
|
||||
"""
|
||||
print(f"[TOOL] filter_data | domain={domain}")
|
||||
|
||||
if filters:
|
||||
print(f"[TOOL] filter_data | filters={filters}")
|
||||
if search:
|
||||
print(f"[TOOL] filter_data | search={search}")
|
||||
|
||||
if domain not in _GLOBAL_RAW_CACHE:
|
||||
return f"⚠️ 请先调用 query_data 查询 {domain} 数据"
|
||||
|
||||
cfg = get_domain(domain)
|
||||
df = pd.DataFrame(_GLOBAL_RAW_CACHE[domain])
|
||||
original_len = len(df)
|
||||
|
||||
# ---------- 筛选(AND) ----------
|
||||
if filters:
|
||||
for f in filters:
|
||||
field, op, value = f["field"], f["op"], f["value"]
|
||||
if field not in df.columns:
|
||||
continue
|
||||
|
||||
col = df[field]
|
||||
# 处理多值情况
|
||||
if col.apply(lambda x: isinstance(x, (list, tuple, set))).any():
|
||||
if op == "contains":
|
||||
col = col.apply(lambda x: value in x if isinstance(x, (list, tuple, set)) else False)
|
||||
elif op == "not_contains":
|
||||
col = col.apply(lambda x: value not in x if isinstance(x, (list, tuple, set)) else True)
|
||||
df = df[col]
|
||||
continue
|
||||
|
||||
# 单值转为字符串比较
|
||||
col = col.astype(str)
|
||||
if op == "eq":
|
||||
df = df[col == value]
|
||||
elif op == "ne":
|
||||
df = df[col != value]
|
||||
elif op == "contains":
|
||||
df = df[col.str.contains(value, na=False)]
|
||||
elif op == "not_contains":
|
||||
df = df[~col.str.contains(value, na=False)]
|
||||
|
||||
# ---------- 检索(OR) ----------
|
||||
if search:
|
||||
mask = pd.Series(False, index=df.index)
|
||||
for kw in search:
|
||||
for field in cfg.search_fields:
|
||||
if field not in df.columns:
|
||||
continue
|
||||
# 如果包含就表示命中
|
||||
mask |= df[field].astype(str).str.contains(kw, na=False)
|
||||
df = df[mask]
|
||||
|
||||
filtered_len = len(df)
|
||||
_GLOBAL_FILTERED_CACHE[domain] = df.to_dict("records")
|
||||
|
||||
print(f"[TOOL] filter_data | original={original_len} filtered={filtered_len}")
|
||||
|
||||
base = f"筛选完成。原始:{original_len} 条,筛选后:{filtered_len} 条。"
|
||||
|
||||
if filtered_len == 0:
|
||||
return base + "未命中任何记录。"
|
||||
|
||||
if filtered_len <= MAX_CONFIRM_COUNT:
|
||||
return base + "如需查看明细,请调用 inspect_data。"
|
||||
else:
|
||||
return base + "数据量较大,建议进一步缩小筛选范围后再查看明细。"
|
||||
|
||||
|
||||
@tool
|
||||
def inspect_data(domain: str, sort_by: str | None = None, ascending: bool = False, top_n: int | None = None) -> str:
|
||||
"""
|
||||
对筛选后的数据进行排序、抽样并按安全策略展示明细。
|
||||
|
||||
展示策略:
|
||||
- ≤ 20 条:直接展示
|
||||
- 21–50 条:询问用户是否查看
|
||||
- > 50 条:拒绝展示,建议缩小范围
|
||||
|
||||
参数:
|
||||
- domain: 业务域名称
|
||||
- sort_by: 排序字段
|
||||
- ascending: 是否升序(默认 False,即降序)
|
||||
- top_n: 返回条数(可选,不传则使用展示策略)
|
||||
"""
|
||||
print(f"[TOOL] inspect_data | domain={domain}")
|
||||
print(f"[TOOL] inspect_data | sort_by={sort_by}, ascending={ascending}, top_n={top_n}")
|
||||
|
||||
if domain not in _GLOBAL_FILTERED_CACHE:
|
||||
print(f"[TOOL] inspect_data | ERROR: filtered cache missing for {domain}")
|
||||
return f"⚠️ 请先调用 filter_data 或 query_data 确定数据集"
|
||||
|
||||
cfg = get_domain(domain)
|
||||
df = pd.DataFrame(_GLOBAL_FILTERED_CACHE[domain])
|
||||
total_after_filter = len(df)
|
||||
print(f"[TOOL] inspect_data | filtered_rows={total_after_filter}")
|
||||
|
||||
# ---------- 排序 ----------
|
||||
if sort_by:
|
||||
if sort_by not in df.columns:
|
||||
print(f"[TOOL] inspect_data | ERROR: sort_by field={sort_by} not exist")
|
||||
return f"⚠️ 字段 {sort_by} 不存在,无法排序"
|
||||
|
||||
try:
|
||||
df[sort_by] = pd.to_numeric(df[sort_by], errors="raise")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
df = df.sort_values(sort_by, ascending=ascending)
|
||||
print(f"[TOOL] inspect_data | sorted by {sort_by} ascending={ascending}")
|
||||
|
||||
# ---------- 抽样(TopN 优先于展示策略) ----------
|
||||
if top_n:
|
||||
df = df.head(top_n)
|
||||
print(f"[TOOL] inspect_data | sampled top_n={top_n}")
|
||||
else:
|
||||
# 未指定 top_n 时,仍可能受展示策略限制
|
||||
pass
|
||||
|
||||
final_len = len(df)
|
||||
print(f"[TOOL] inspect_data | final_rows={final_len}, show_fields={cfg.detail_fields}")
|
||||
|
||||
if df.empty:
|
||||
return "⚠️ 当前数据集为空,无法展示明细。"
|
||||
|
||||
# ---------- 展示策略(核心) ----------
|
||||
# 情况 1:少量数据,直接展示
|
||||
if final_len <= MAX_DETAIL_COUNT:
|
||||
lines = [f"💡 共展示 {final_len} 条记录"]
|
||||
for _, row in df.iterrows():
|
||||
lines.append("---")
|
||||
for fld in cfg.detail_fields:
|
||||
if fld not in row:
|
||||
continue
|
||||
val = row[fld]
|
||||
if pd.isna(val) or val == "":
|
||||
val = "未填写"
|
||||
lines.append(f"- {cfg.get_chinese_name(fld)}:{val}")
|
||||
print(f"[TOOL] inspect_data | rendered directly")
|
||||
return "\n".join(lines)
|
||||
|
||||
# 情况 2:中等数据量,询问用户
|
||||
if final_len <= MAX_CONFIRM_COUNT:
|
||||
print(f"[TOOL] inspect_data | confirm required")
|
||||
return (
|
||||
f"筛选后共 {final_len} 条记录,是否查看明细?\n"
|
||||
f"(提示:可使用 top_n 参数只查看前 N 条,如最大的一笔)"
|
||||
)
|
||||
|
||||
# 情况 3:数据量过大,拒绝展示
|
||||
print(f"[TOOL] inspect_data | too large, rejected")
|
||||
return (
|
||||
f"数据量过大({final_len} 条),暂不展示明细。\n"
|
||||
f"建议:缩小筛选范围;\n"
|
||||
)
|
||||
|
||||
|
||||
@tool
|
||||
def analyze_data(domain: str, dimensions: list[str], metrics: list[str] | None = None) -> str:
|
||||
"""
|
||||
对数据进行多维度统计分析。
|
||||
|
||||
参数:
|
||||
- domain: 业务域名称
|
||||
- dimensions: 英文维度字段列表(如 ["category", "merchant"])
|
||||
- metrics: 指标字段列表(如 ["amount"])
|
||||
若不传,仅统计条数
|
||||
|
||||
示例:
|
||||
{
|
||||
"domain": "bill",
|
||||
"dimensions": ["category"],
|
||||
"metrics": ["amount"]
|
||||
}
|
||||
"""
|
||||
print(f"[TOOL] analyze_data | domain={domain} | dimensions={dimensions} | metrics={metrics}")
|
||||
|
||||
if domain not in _GLOBAL_FILTERED_CACHE:
|
||||
return f"⚠️ 请先调用 query_data 查询 {domain} 数据"
|
||||
|
||||
cfg = get_domain(domain)
|
||||
df = pd.DataFrame(_GLOBAL_FILTERED_CACHE[domain])
|
||||
|
||||
# ---------- 基础信息 ----------
|
||||
lines = [
|
||||
"统计完成。",
|
||||
f"数据总量:{len(df)} 条",
|
||||
f"统计维度:{', '.join(cfg.get_chinese_name(d) for d in dimensions)}",
|
||||
]
|
||||
|
||||
metrics = metrics or []
|
||||
metric_meta = {m.field: m for m in cfg.metrics if m.field in metrics}
|
||||
|
||||
if metric_meta:
|
||||
metric_names = ", ".join(m.name for m in metric_meta.values())
|
||||
lines.append(f"统计指标:条数、{metric_names}\n")
|
||||
else:
|
||||
lines.append("")
|
||||
|
||||
# ---------- 按维度统计 ----------
|
||||
for field in dimensions:
|
||||
if field not in df.columns:
|
||||
continue
|
||||
|
||||
col = df[field]
|
||||
|
||||
# 多值字段拆行
|
||||
if col.apply(lambda x: isinstance(x, (list, tuple, set))).any():
|
||||
print(f"[WARN] analyze_data | field={field} is multi-value, metrics may be duplicated")
|
||||
df_tmp = df.explode(field)
|
||||
else:
|
||||
df_tmp = df
|
||||
|
||||
# 构建聚合规则
|
||||
agg_dict = {
|
||||
"数量": (cfg.primary_key, "count")
|
||||
}
|
||||
|
||||
for m in metrics:
|
||||
meta = cfg.get_metric(m)
|
||||
if not meta or meta.field not in df_tmp.columns:
|
||||
continue
|
||||
agg_dict[meta.name] = (meta.field, meta.agg)
|
||||
|
||||
if not agg_dict:
|
||||
continue
|
||||
|
||||
agg = (
|
||||
df_tmp.groupby(field, dropna=False)
|
||||
.agg(**agg_dict)
|
||||
.sort_values("数量", ascending=False)
|
||||
.reset_index()
|
||||
)
|
||||
|
||||
lines.append(f"【{cfg.get_chinese_name(field)}】")
|
||||
|
||||
for _, row in agg.iterrows():
|
||||
parts = [f"{row[field]}:{int(row['数量'])} 条"]
|
||||
|
||||
for m in metrics:
|
||||
meta = cfg.get_metric(m)
|
||||
if not meta:
|
||||
continue
|
||||
|
||||
col_name = meta.name
|
||||
if col_name not in row:
|
||||
continue
|
||||
|
||||
val = row[col_name]
|
||||
parts.append(
|
||||
f"{meta.name} {format_value(val, meta.format)}"
|
||||
)
|
||||
|
||||
lines.append(",".join(parts))
|
||||
|
||||
lines.append("")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def format_value(val, fmt: str) -> str:
|
||||
if pd.isna(val):
|
||||
return "—"
|
||||
|
||||
if fmt == "money":
|
||||
return f"¥{val:,.2f}"
|
||||
if fmt == "duration":
|
||||
return f"{val:.1f} 分钟"
|
||||
if fmt == "percent":
|
||||
return f"{val * 100:.1f}%"
|
||||
if fmt.startswith("number:"):
|
||||
prec = int(fmt.split(":")[1])
|
||||
return f"{val:.{prec}f}"
|
||||
if fmt == "int":
|
||||
return f"{int(val)}"
|
||||
|
||||
# auto / fallback
|
||||
return str(val)
|
||||
Reference in New Issue
Block a user