413 lines
16 KiB
Python
413 lines
16 KiB
Python
from __future__ import annotations
|
||
|
||
"""把线上 ELK case 直接转换为 product-data canonical records。
|
||
|
||
这个脚本承接 online-mining-v2 和 product-data:
|
||
1. 按 request_id 拉主 NLP 表和前处理表。
|
||
2. 优先从前处理表 promptModel 的 [对话历史] 抽取模型真实输入 session。
|
||
3. 从用户指定 target 或线上 planning/code 结果推断标签。
|
||
4. 直接产出 canonical records,并可同步导出流通表格、训练 jsonl、评测 csv。
|
||
"""
|
||
|
||
import json
|
||
import re
|
||
import sys
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
from online_mining_common import (
|
||
OnlineMiningError,
|
||
compact_text,
|
||
emit_error,
|
||
emit_success,
|
||
extract_case,
|
||
fetch_one_by_request_id,
|
||
load_json_payload,
|
||
resolve_portable_path,
|
||
string_values,
|
||
)
|
||
|
||
PRODUCT_DATA_SCRIPT_DIR = Path(__file__).resolve().parents[1].parent / "product-data" / "scripts"
|
||
sys.path.insert(0, str(PRODUCT_DATA_SCRIPT_DIR))
|
||
|
||
from product_data_portable import ( # noqa: E402
|
||
canonical_target,
|
||
export_dataset_records,
|
||
export_planning_eval_csv,
|
||
export_training_jsonl,
|
||
record_id,
|
||
target_type,
|
||
validate_dataset_records,
|
||
)
|
||
|
||
|
||
def load_cases(payload: dict[str, Any]) -> list[dict[str, Any]]:
|
||
if isinstance(payload.get("cases"), list):
|
||
return [item for item in payload["cases"] if isinstance(item, dict)]
|
||
cases_path = payload.get("cases_path")
|
||
if cases_path:
|
||
rows = []
|
||
text = resolve_portable_path(str(cases_path)).read_text(encoding="utf-8")
|
||
for line in text.splitlines():
|
||
if not line.strip():
|
||
continue
|
||
item = json.loads(line)
|
||
if isinstance(item, dict):
|
||
rows.append(item)
|
||
return rows
|
||
return []
|
||
|
||
|
||
def merge_case(request_id: str, main_case: dict[str, Any] | None, pre_case: dict[str, Any] | None) -> dict[str, Any]:
|
||
return {
|
||
"request_id": request_id,
|
||
"query": (pre_case or {}).get("query") or (main_case or {}).get("query"),
|
||
"timestamp": (main_case or {}).get("timestamp") or (pre_case or {}).get("timestamp"),
|
||
"session_id": (main_case or {}).get("session_id"),
|
||
"device_id": (main_case or {}).get("device_id"),
|
||
"device": (main_case or {}).get("device"),
|
||
"tts": (main_case or {}).get("text"),
|
||
"main": main_case,
|
||
"pre_processing": pre_case,
|
||
}
|
||
|
||
|
||
def fetch_cases_by_request_ids(
|
||
request_ids: list[str],
|
||
*,
|
||
date: str | None,
|
||
date_from: str | None = None,
|
||
date_to: str | None = None,
|
||
lookback_days: int | None = None,
|
||
) -> list[dict[str, Any]]:
|
||
rows = []
|
||
for request_id in request_ids:
|
||
main_doc = fetch_one_by_request_id(
|
||
"main",
|
||
request_id,
|
||
date,
|
||
date_from=date_from,
|
||
date_to=date_to,
|
||
lookback_days=lookback_days,
|
||
)
|
||
pre_doc = fetch_one_by_request_id(
|
||
"pre_processing",
|
||
request_id,
|
||
date,
|
||
date_from=date_from,
|
||
date_to=date_to,
|
||
lookback_days=lookback_days,
|
||
)
|
||
main_case = extract_case("main", main_doc) if main_doc else None
|
||
pre_case = extract_case("pre_processing", pre_doc) if pre_doc else None
|
||
rows.append(merge_case(request_id, main_case, pre_case))
|
||
return rows
|
||
|
||
|
||
def case_value(case: dict[str, Any], key: str) -> Any:
|
||
if key in case:
|
||
return case.get(key)
|
||
main = case.get("main") if isinstance(case.get("main"), dict) else {}
|
||
pre = case.get("pre_processing") if isinstance(case.get("pre_processing"), dict) else {}
|
||
return case.get(key) or pre.get(key) or main.get(key)
|
||
|
||
|
||
def enrich_prev_session(case: dict[str, Any], *, limit: int) -> list[dict[str, Any]]:
|
||
explicit = case.get("prev_session")
|
||
if isinstance(explicit, list):
|
||
return normalize_prev_session(explicit)[-limit:]
|
||
timestamp = optional_int(case_value(case, "timestamp"))
|
||
# 前处理 promptModel 是中控 planning 模型真实看到的输入,里面的 [对话历史]
|
||
# 才是训练/评测应复现的 session。主表 session_id 不是同一概念,不默认使用。
|
||
prompt_history = parse_prompt_history(str(case_value(case, "prompt_model") or ""))
|
||
if prompt_history and timestamp is not None:
|
||
step_ms = 60_000
|
||
start_ts = timestamp - len(prompt_history) * step_ms
|
||
for index, item in enumerate(prompt_history):
|
||
item["timestamp"] = start_ts + index * step_ms
|
||
return prompt_history[-limit:]
|
||
|
||
|
||
def parse_prompt_history(prompt: str) -> list[dict[str, Any]]:
|
||
if not prompt:
|
||
return []
|
||
history_marker = "[对话历史]"
|
||
current_marker = "[当前query]"
|
||
history_start = prompt.rfind(history_marker)
|
||
if history_start < 0:
|
||
return []
|
||
current_start = prompt.find(current_marker, history_start)
|
||
if current_start < 0:
|
||
return []
|
||
block = prompt[history_start + len(history_marker) : current_start].strip()
|
||
if not block:
|
||
return []
|
||
rows: list[dict[str, Any]] = []
|
||
pending_user: str | None = None
|
||
for raw_line in block.splitlines():
|
||
line = raw_line.strip()
|
||
if not line:
|
||
continue
|
||
user_match = re.match(r"^用户\s*[::]\s*(.*)$", line)
|
||
if user_match:
|
||
if pending_user is not None:
|
||
rows.append({"query": pending_user, "tts": "", "timestamp": 0})
|
||
pending_user = user_match.group(1).strip()
|
||
continue
|
||
assistant_match = re.match(r"^小爱\s*[::]\s*(.*)$", line)
|
||
if assistant_match and pending_user is not None:
|
||
rows.append({"query": pending_user, "tts": assistant_match.group(1).strip(), "timestamp": 0})
|
||
pending_user = None
|
||
if pending_user is not None:
|
||
rows.append({"query": pending_user, "tts": "", "timestamp": 0})
|
||
return [row for row in rows if row["query"]]
|
||
|
||
|
||
def normalize_prev_session(items: list[Any]) -> list[dict[str, Any]]:
|
||
rows = []
|
||
for item in items:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
query = str(item.get("query") or "").strip()
|
||
if not query:
|
||
continue
|
||
rows.append(
|
||
{
|
||
"query": query,
|
||
"tts": str(item.get("tts") or ""),
|
||
"timestamp": optional_int(item.get("timestamp")) or 0,
|
||
}
|
||
)
|
||
rows.sort(key=lambda item: int(item.get("timestamp") or 0))
|
||
return rows
|
||
|
||
|
||
def infer_target(case: dict[str, Any], explicit_target: str) -> tuple[str, str]:
|
||
if explicit_target.strip():
|
||
return canonical_target(explicit_target.strip()), "input.target"
|
||
case_target = str(case_value(case, "target") or "").strip()
|
||
if case_target:
|
||
return canonical_target(case_target), "case.target"
|
||
|
||
pre = case.get("pre_processing") if isinstance(case.get("pre_processing"), dict) else {}
|
||
main = case.get("main") if isinstance(case.get("main"), dict) else {}
|
||
for key in ("code", "planning_result"):
|
||
value = str(pre.get(key) or "").strip()
|
||
if value:
|
||
return canonical_target(value), f"pre_processing.{key}"
|
||
|
||
llm_agent_info = main.get("llm_agent_info")
|
||
if isinstance(llm_agent_info, dict):
|
||
agent_type = str(llm_agent_info.get("agentType") or "").strip()
|
||
if agent_type:
|
||
return f'Agent(tag="{agent_type}")', "main.intention.intent_arbitrator_info.llm_agent_info.agentType"
|
||
|
||
domain = str(main.get("domain") or case_value(case, "domain") or "").strip()
|
||
if domain:
|
||
if re.fullmatch(r"[A-Z][A-Za-z0-9_]*", domain):
|
||
return f"{domain}()", "main.domain"
|
||
return f'Agent(tag="{domain}")', "main.domain"
|
||
raise OnlineMiningError("target is required because online logs do not contain code/planning_result/domain")
|
||
|
||
|
||
def case_to_record(
|
||
case: dict[str, Any],
|
||
*,
|
||
dataset_label: str,
|
||
target: str,
|
||
complex_value: bool,
|
||
index: int,
|
||
session_limit: int,
|
||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||
query = str(case_value(case, "query") or "").strip()
|
||
if not query:
|
||
raise OnlineMiningError(f"case {index} has no query")
|
||
request_id = str(case_value(case, "request_id") or "").strip()
|
||
if not request_id:
|
||
raise OnlineMiningError(f"case {index} has no request_id")
|
||
timestamp = optional_int(case_value(case, "timestamp"))
|
||
if timestamp is None:
|
||
raise OnlineMiningError(f"case {index} has no timestamp")
|
||
|
||
final_target, target_source = infer_target(case, target)
|
||
prev_session = enrich_prev_session(case, limit=session_limit)
|
||
main = case.get("main") if isinstance(case.get("main"), dict) else {}
|
||
pre = case.get("pre_processing") if isinstance(case.get("pre_processing"), dict) else {}
|
||
record = {
|
||
"record_id": record_id(source_type="online", batch_id="online", request_id=request_id, index=index),
|
||
"source": {
|
||
"type": "online",
|
||
"request_id": request_id,
|
||
"timestamp": timestamp,
|
||
},
|
||
"turn": {
|
||
"query": query,
|
||
"timestamp": timestamp,
|
||
},
|
||
"prev_session": prev_session,
|
||
"context": {},
|
||
"label": {
|
||
"dataset_label": dataset_label,
|
||
"target": final_target,
|
||
"target_type": target_type(final_target),
|
||
},
|
||
"dimensions": {
|
||
"complex": parse_complex(case_value(case, "complex"), default=complex_value),
|
||
},
|
||
"meta": {
|
||
"session_id": str(case_value(case, "session_id") or ""),
|
||
"device_id": str(case_value(case, "device_id") or ""),
|
||
"device": compact_text(case_value(case, "device"), 500),
|
||
"online_tts": str(case_value(case, "tts") or ""),
|
||
"main_domain": str(main.get("domain") or ""),
|
||
"main_func": str(main.get("func") or ""),
|
||
"planning_result": str(pre.get("planning_result") or ""),
|
||
"planning_code": str(pre.get("code") or ""),
|
||
"hit_rules": pre.get("hit_rules") or main.get("hit_rules") or [],
|
||
"candidate_domains": pre.get("candidate_domains") or [],
|
||
"target_source": target_source,
|
||
},
|
||
}
|
||
summary = {
|
||
"request_id": request_id,
|
||
"query": query,
|
||
"target": final_target,
|
||
"target_source": target_source,
|
||
"prev_session_count": len(prev_session),
|
||
"main_domain": record["meta"]["main_domain"],
|
||
"planning_result": record["meta"]["planning_result"],
|
||
}
|
||
return record, summary
|
||
|
||
|
||
def optional_int(value: Any) -> int | None:
|
||
if value is None or value == "":
|
||
return None
|
||
try:
|
||
return int(value)
|
||
except (TypeError, ValueError):
|
||
return None
|
||
|
||
|
||
def parse_complex(value: Any, *, default: bool) -> bool:
|
||
if value is None or value == "":
|
||
return default
|
||
if isinstance(value, bool):
|
||
return value
|
||
text = str(value).strip().lower()
|
||
if text in {"true", "1", "yes", "y", "是", "复杂", "complex"}:
|
||
return True
|
||
if text in {"false", "0", "no", "n", "否", "不复杂", "简单", "simple"}:
|
||
return False
|
||
raise OnlineMiningError("complex must be true/false")
|
||
|
||
|
||
def write_records_jsonl(records: list[dict[str, Any]], path: str) -> str:
|
||
output = resolve_portable_path(path)
|
||
output.parent.mkdir(parents=True, exist_ok=True)
|
||
with output.open("w", encoding="utf-8") as fh:
|
||
for record in records:
|
||
fh.write(json.dumps(record, ensure_ascii=False, separators=(",", ":")) + "\n")
|
||
return str(output)
|
||
|
||
|
||
def main() -> int:
|
||
try:
|
||
payload = load_json_payload()
|
||
dataset_label = str(payload.get("dataset_label") or "").strip()
|
||
if not dataset_label:
|
||
raise OnlineMiningError("dataset_label is required")
|
||
target = str(payload.get("target") or "").strip()
|
||
complex_value = parse_complex(payload.get("complex"), default=False)
|
||
date = str(payload.get("date") or "").strip() or None
|
||
date_from = str(payload.get("date_from") or "").strip() or None
|
||
date_to = str(payload.get("date_to") or "").strip() or None
|
||
lookback_days = payload.get("lookback_days")
|
||
lookback_days = int(lookback_days) if lookback_days not in (None, "") else None
|
||
session_limit = int(payload.get("session_limit") or 10)
|
||
if session_limit < 0 or session_limit > 10:
|
||
raise OnlineMiningError("session_limit must be between 0 and 10")
|
||
|
||
request_ids = string_values(payload.get("request_ids"))
|
||
cases = load_cases(payload)
|
||
if request_ids:
|
||
cases = fetch_cases_by_request_ids(
|
||
request_ids,
|
||
date=date,
|
||
date_from=date_from,
|
||
date_to=date_to,
|
||
lookback_days=lookback_days,
|
||
)
|
||
if not cases:
|
||
raise OnlineMiningError("request_ids, cases or cases_path is required")
|
||
|
||
records = []
|
||
summaries = []
|
||
skipped = []
|
||
for index, case in enumerate(cases, start=1):
|
||
try:
|
||
record, summary = case_to_record(
|
||
case,
|
||
dataset_label=dataset_label,
|
||
target=target,
|
||
complex_value=complex_value,
|
||
index=index,
|
||
session_limit=session_limit,
|
||
)
|
||
except Exception as exc: # noqa: BLE001
|
||
skipped.append({"index": index, "request_id": case.get("request_id"), "error": str(exc)})
|
||
continue
|
||
records.append(record)
|
||
summaries.append(summary)
|
||
|
||
if not records:
|
||
raise OnlineMiningError(f"no records built; skipped={skipped}")
|
||
|
||
records_path = write_records_jsonl(records, str(payload.get("records_output_path") or "scratchpad/online_records.jsonl"))
|
||
validation = validate_dataset_records(records)
|
||
output_dir = str(payload.get("output_dir") or "output")
|
||
exports: dict[str, Any] = {}
|
||
if bool(payload.get("export_records", True)):
|
||
exports["records"] = export_dataset_records(
|
||
records,
|
||
output_path=str(resolve_portable_path(output_dir) / "records.jsonl"),
|
||
output_format="jsonl",
|
||
require_validation_ok=bool(payload.get("require_validation_ok", True)),
|
||
overwrite=bool(payload.get("overwrite", True)),
|
||
export_table=bool(payload.get("export_table", True)),
|
||
)
|
||
if bool(payload.get("export_training", False)):
|
||
exports["training"] = export_training_jsonl(
|
||
records,
|
||
output_path=str(resolve_portable_path(output_dir) / "training.jsonl"),
|
||
require_validation_ok=bool(payload.get("require_validation_ok", True)),
|
||
overwrite=bool(payload.get("overwrite", True)),
|
||
)
|
||
if bool(payload.get("export_eval", False)):
|
||
exports["eval"] = export_planning_eval_csv(
|
||
records,
|
||
output_path=str(resolve_portable_path(output_dir) / "eval_planning.csv"),
|
||
require_validation_ok=bool(payload.get("require_validation_ok", True)),
|
||
overwrite=bool(payload.get("overwrite", True)),
|
||
)
|
||
|
||
emit_success(
|
||
{
|
||
"date": date or (f"{date_from}..{date_to}" if date_from and date_to else f"past-{lookback_days}d" if lookback_days else "past-48h"),
|
||
"record_count": len(records),
|
||
"records_path": records_path,
|
||
"validation": validation,
|
||
"exports": exports,
|
||
"summaries": summaries,
|
||
"skipped": skipped,
|
||
}
|
||
)
|
||
return 0
|
||
except Exception as exc: # noqa: BLE001
|
||
emit_error(exc)
|
||
return 1
|
||
|
||
|
||
if __name__ == "__main__":
|
||
raise SystemExit(main())
|