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)