Files
sweet-hut-agent/tools/base.py

369 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import numpy as np
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]
# 处理多值情况
is_multi = col.apply(lambda x: isinstance(x, (list, tuple, set))).any()
if is_multi:
if op == "eq":
col = col.apply(lambda x: value in x if isinstance(x, (list, tuple, set)) else False)
elif op == "ne":
col = col.apply(lambda x: value not in x if isinstance(x, (list, tuple, set)) else True)
elif 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)
else:
continue
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
# list 转成可检索的字符串
col_str = df[field].apply(
lambda x: " ".join(map(str, x))
if isinstance(x, (list, tuple, set))
else str(x)
)
# 如果包含则命中
mask |= col_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 条:直接展示
- 2150 条:询问用户是否查看
- > 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]
# 数组 / Series
if isinstance(val, (np.ndarray, pd.Series)) or (isinstance(val, (list, tuple))):
if len(val) == 0:
val = "未填写"
else:
val = "".join(map(str, val))
elif pd.isna(val):
val = "未填写"
elif val == "":
val = "未填写"
else:
val = str(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)