feat:更新智能体工具

This commit is contained in:
2026-07-13 22:06:05 +08:00
parent 900d5bd774
commit 5213897bbd
14 changed files with 672 additions and 157 deletions

338
tools/base.py Normal file
View 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 条:直接展示
- 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]
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)

View File

@@ -1,67 +0,0 @@
from langchain_core.tools import tool
import requests
BASE_URL = "https://cxx0822.s.3q.hair/home-api/"
USERNAME = "Cxx0822"
PASSWORD = "19940822Cxx"
def get_token() -> str:
"""
调用登录接口,返回 saToken.tokenValue
"""
params = {
"name": USERNAME,
"password": PASSWORD
}
try:
resp = requests.post(f"{BASE_URL}session", params=params, timeout=10)
resp.raise_for_status()
data = resp.json()
# 按你给的返回结构取值
token_value = data["saToken"]["tokenValue"]
return token_value
except requests.exceptions.RequestException as e:
raise RuntimeError(f"登录失败,无法获取 token: {e}")
except KeyError as e:
raise RuntimeError(f"登录响应结构异常,未找到 tokenValue: {e}")
@tool
def bill_search(start_date: str, end_date: str) -> str:
"""
查询账单列表。
参数说明:
- start_date (必填):
- 开始日期格式YYYY-MM-DD
- 示例2026-06-01
- end_date (必填):
- 结束日期格式YYYY-MM-DD
- 示例2026-07-01
"""
print(f"start_date: {start_date}, end_date: {end_date}")
params = {
"bookName": "",
"type": "",
"startDate": start_date,
"endDate": end_date
}
token = get_token()
headers = {
"satoken": token
}
try:
resp = requests.get(f"{BASE_URL}bill/summary", params=params, headers=headers, timeout=10)
resp.raise_for_status()
return resp.text
except Exception as e:
return f"查询账单失败: {e}"