Files
zk-data-agent/src/data_agent_inputs.py
T
2026-05-08 17:09:07 +08:00

619 lines
23 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.
from __future__ import annotations
"""数据智能体输入文件解析工具。
这里的工具只负责把产品定义、badcase 表格、走查文档等输入稳定抽成
结构化证据。语义判断和生成策略仍交给 skill/模型完成,避免工具里写死业务规则。
"""
import csv
import json
import re
import zipfile
from pathlib import Path
from typing import Any
from xml.etree import ElementTree
class DataAgentInputError(ValueError):
"""输入文件无法解析或参数结构不符合预期。"""
SUPPORTED_SUFFIXES = {'.csv', '.docx', '.json', '.jsonl', '.md', '.pdf', '.txt', '.xlsx'}
QUERY_ROLE = 'query'
EXPECTED_LABEL_ROLE = 'expected_label'
PREDICTED_LABEL_ROLE = 'predicted_label'
CONTEXT_ROLE = 'context'
NOTES_ROLE = 'notes'
CASE_TYPE_ROLE = 'case_type'
def load_input_sources(
root: str | Path,
paths: list[str],
*,
max_files: int = 20,
max_paragraphs_per_file: int = 80,
max_tables_per_file: int = 20,
max_rows_per_table: int = 30,
max_cell_chars: int = 240,
) -> dict[str, Any]:
resolved_paths = _resolve_input_paths(root, paths, max_files=max_files)
sources = [
_load_one_source(
path,
root=Path(root).resolve(),
max_paragraphs=max_paragraphs_per_file,
max_tables=max_tables_per_file,
max_rows=max_rows_per_table,
max_cell_chars=max_cell_chars,
)
for path in resolved_paths
]
return {'sources': sources, 'source_count': len(sources)}
def extract_case_evidence(
root: str | Path,
*,
paths: list[str] | None = None,
loaded_sources: dict[str, Any] | None = None,
field_mapping: dict[str, str] | None = None,
max_cases: int = 200,
) -> dict[str, Any]:
sources = _sources_from_args(root, paths, loaded_sources)
evidence: list[dict[str, Any]] = []
table_profiles: list[dict[str, Any]] = []
required_questions: list[str] = []
for source in sources:
for table in source.get('tables', []):
headers = _table_headers(table)
header_index = _table_header_index(table)
if not headers:
continue
rows = table.get('rows')
if not isinstance(rows, list):
continue
profile = _profile_headers(headers, rows, field_mapping or {})
profile.update({'source_ref': table.get('source_ref'), 'title': table.get('title', '')})
table_profiles.append(profile)
questions = _questions_for_profile(profile)
required_questions.extend(questions)
if questions:
continue
evidence.extend(
_rows_to_case_evidence(
source,
table,
rows,
profile['role_to_column'],
header_index=header_index,
max_cases=max(0, max_cases - len(evidence)),
)
)
if len(evidence) >= max_cases:
break
if len(evidence) >= max_cases:
break
return {
'evidence': evidence,
'evidence_count': len(evidence),
'table_profiles': table_profiles,
'required_questions': _dedupe(required_questions),
}
def render_source_context(
root: str | Path,
*,
paths: list[str] | None = None,
loaded_sources: dict[str, Any] | None = None,
max_chars: int = 40_000,
max_tables_per_source: int = 12,
max_rows_per_table: int = 20,
focus_keywords: list[str] | None = None,
) -> dict[str, Any]:
sources = _sources_from_args(root, paths, loaded_sources)
sections: list[str] = []
truncated = False
keywords = [keyword for keyword in focus_keywords or [] if isinstance(keyword, str) and keyword.strip()]
for source in sources:
source_lines = [
f'# Source: {source.get("path", "")}',
f'kind: {source.get("kind", "")}',
f'summary: {source.get("summary", "")}',
]
warnings = source.get('warnings')
if warnings:
source_lines.append('warnings: ' + '; '.join(str(item) for item in warnings))
for paragraph in source.get('paragraphs', []):
text = str(paragraph.get('text') or '').strip()
if not text or not _matches_focus(text, keywords):
continue
source_lines.append('')
source_lines.append(f'## Paragraph {paragraph.get("source_ref", "")}')
source_lines.append(text)
for table in source.get('tables', [])[:max_tables_per_source]:
headers = _table_headers(table)
rows = table.get('rows')
if not isinstance(rows, list):
continue
rendered_rows = []
for row_index, row in enumerate(rows[:max_rows_per_table], start=1):
if not isinstance(row, list):
continue
row_text = _render_table_row(headers, row) if headers else ' | '.join(str(cell) for cell in row)
if not row_text or not _matches_focus(row_text, keywords):
continue
rendered_rows.append(f'[{table.get("source_ref")}:row{row_index}] {row_text}')
if rendered_rows:
source_lines.append('')
source_lines.append(f'## Table {table.get("title", "")} {table.get("source_ref", "")}')
if headers:
source_lines.append('columns: ' + ' | '.join(headers))
source_lines.extend(rendered_rows)
section = '\n'.join(source_lines).strip()
if section:
sections.append(section)
rendered = '\n\n'.join(sections)
if len(rendered) > max_chars:
rendered = rendered[: max_chars - 80].rstrip() + '\n\n...[truncated: source context exceeded max_chars]...'
truncated = True
return {
'context_text': rendered,
'char_count': len(rendered),
'truncated': truncated,
'source_count': len(sources),
'instructions': '请让大模型只基于 context_text 抽取和用户 query 语义、功能点、边界、target/open questions 相关的信息,并保留 source refs。',
}
def _resolve_input_paths(root: str | Path, paths: list[str], *, max_files: int) -> list[Path]:
if not isinstance(paths, list) or not paths:
raise DataAgentInputError('paths must be a non-empty array')
root_path = Path(root).resolve()
resolved: list[Path] = []
for raw in paths:
if not isinstance(raw, str) or not raw.strip():
raise DataAgentInputError('paths items must be non-empty strings')
path = Path(raw).expanduser()
candidate = path if path.is_absolute() else root_path / path
candidate = candidate.resolve()
if not candidate.exists():
raise DataAgentInputError(f'path not found: {raw}')
if candidate.is_dir():
for child in sorted(candidate.rglob('*')):
if child.is_file() and child.suffix.lower() in SUPPORTED_SUFFIXES:
resolved.append(child.resolve())
elif candidate.is_file():
if candidate.suffix.lower() in SUPPORTED_SUFFIXES:
resolved.append(candidate)
if len(resolved) >= max_files:
break
return _dedupe_paths(resolved)[:max_files]
def _load_one_source(
path: Path,
*,
root: Path,
max_paragraphs: int,
max_tables: int,
max_rows: int,
max_cell_chars: int,
) -> dict[str, Any]:
suffix = path.suffix.lower()
rel_path = _display_input_path(path, root)
source: dict[str, Any] = {
'path': rel_path,
'kind': suffix.lstrip('.'),
'size_bytes': path.stat().st_size,
'paragraphs': [],
'tables': [],
'warnings': [],
}
try:
if suffix == '.xlsx':
_load_xlsx(path, source, max_tables=max_tables, max_rows=max_rows, max_cell_chars=max_cell_chars)
elif suffix == '.csv':
_load_csv(path, source, max_rows=max_rows, max_cell_chars=max_cell_chars)
elif suffix == '.docx':
_load_docx(path, source, max_paragraphs=max_paragraphs, max_tables=max_tables, max_rows=max_rows, max_cell_chars=max_cell_chars)
elif suffix == '.pdf':
_load_pdf(path, source, max_paragraphs=max_paragraphs)
elif suffix in {'.txt', '.md', '.json', '.jsonl'}:
_load_text(path, source, max_paragraphs=max_paragraphs)
else:
source['warnings'].append(f'暂不支持的文件类型: {suffix}')
except Exception as exc:
source['warnings'].append(f'解析失败: {type(exc).__name__}: {exc}')
source['summary'] = _source_summary(source)
return source
def _display_input_path(path: Path, root: Path) -> str:
try:
return path.relative_to(root).as_posix()
except ValueError:
return path.as_posix()
def _load_xlsx(path: Path, source: dict[str, Any], *, max_tables: int, max_rows: int, max_cell_chars: int) -> None:
try:
import openpyxl # type: ignore[import-not-found]
except ImportError as exc:
raise DataAgentInputError('解析 xlsx 需要 openpyxl') from exc
workbook = openpyxl.load_workbook(path, data_only=True, read_only=True)
for sheet_index, sheet in enumerate(workbook.worksheets[:max_tables]):
# 有些 xlsx 的维度元数据不完整,read_only 模式会误判成 A1:A1。
# reset_dimensions 后让 openpyxl 从实际行流里重新推断列数。
if hasattr(sheet, 'reset_dimensions'):
sheet.reset_dimensions()
rows = []
for row in sheet.iter_rows(max_row=max_rows, values_only=True):
rows.append([_clean_cell(value, max_cell_chars=max_cell_chars) for value in row])
source['tables'].append(
{
'source_ref': f'{source["path"]}#sheet={sheet.title}',
'title': sheet.title,
'row_count': sheet.max_row,
'column_count': sheet.max_column,
'rows': rows,
}
)
if sheet_index + 1 >= max_tables:
break
def _load_csv(path: Path, source: dict[str, Any], *, max_rows: int, max_cell_chars: int) -> None:
rows: list[list[str]] = []
with path.open('r', encoding='utf-8-sig', errors='replace', newline='') as handle:
reader = csv.reader(handle)
for index, row in enumerate(reader):
if index >= max_rows:
break
rows.append([_clean_cell(cell, max_cell_chars=max_cell_chars) for cell in row])
source['tables'].append(
{
'source_ref': source['path'],
'title': path.name,
'row_count': len(rows),
'column_count': max((len(row) for row in rows), default=0),
'rows': rows,
}
)
def _load_docx(
path: Path,
source: dict[str, Any],
*,
max_paragraphs: int,
max_tables: int,
max_rows: int,
max_cell_chars: int,
) -> None:
with zipfile.ZipFile(path) as archive:
document_xml = archive.read('word/document.xml')
root = ElementTree.fromstring(document_xml)
namespace = {'w': 'http://schemas.openxmlformats.org/wordprocessingml/2006/main'}
paragraphs: list[dict[str, str]] = []
tables: list[dict[str, Any]] = []
body = root.find('w:body', namespace)
if body is None:
return
for element in list(body):
if element.tag.endswith('}p') and len(paragraphs) < max_paragraphs:
text = _docx_text(element, namespace)
if text:
paragraphs.append({'text': text, 'source_ref': f'{source["path"]}#p{len(paragraphs) + 1}'})
elif element.tag.endswith('}tbl') and len(tables) < max_tables:
rows = _docx_table_rows(element, namespace, max_rows=max_rows, max_cell_chars=max_cell_chars)
tables.append(
{
'source_ref': f'{source["path"]}#table={len(tables)}',
'title': f'table_{len(tables)}',
'row_count': len(rows),
'column_count': max((len(row) for row in rows), default=0),
'rows': rows,
}
)
source['paragraphs'].extend(paragraphs)
source['tables'].extend(tables)
def _load_pdf(path: Path, source: dict[str, Any], *, max_paragraphs: int) -> None:
try:
from pypdf import PdfReader # type: ignore[import-not-found]
except ImportError:
source['warnings'].append('当前 Python 环境缺少 pypdf,无法抽取 PDF 文本。优先使用同名 docx。')
return
reader = PdfReader(str(path))
for index, page in enumerate(reader.pages[:max_paragraphs]):
text = ' '.join((page.extract_text() or '').split())
if text:
source['paragraphs'].append({'text': text, 'source_ref': f'{source["path"]}#page={index + 1}'})
def _load_text(path: Path, source: dict[str, Any], *, max_paragraphs: int) -> None:
text = path.read_text(encoding='utf-8', errors='replace')
for index, block in enumerate(re.split(r'\n\s*\n', text)):
if index >= max_paragraphs:
break
block = ' '.join(block.split())
if block:
source['paragraphs'].append({'text': block, 'source_ref': f'{source["path"]}#block={index + 1}'})
def _docx_text(element: ElementTree.Element, namespace: dict[str, str]) -> str:
parts = [node.text or '' for node in element.findall('.//w:t', namespace)]
return ' '.join(''.join(parts).split())
def _docx_table_rows(
table: ElementTree.Element,
namespace: dict[str, str],
*,
max_rows: int,
max_cell_chars: int,
) -> list[list[str]]:
rows: list[list[str]] = []
for row in table.findall('./w:tr', namespace):
if len(rows) >= max_rows:
break
cells = []
for cell in row.findall('./w:tc', namespace):
cells.append(_clean_cell(_docx_text(cell, namespace), max_cell_chars=max_cell_chars))
rows.append(cells)
return rows
def _sources_from_args(
root: str | Path,
paths: list[str] | None,
loaded_sources: dict[str, Any] | None,
) -> list[dict[str, Any]]:
if loaded_sources is not None:
sources = loaded_sources.get('sources') if isinstance(loaded_sources, dict) else None
if not isinstance(sources, list):
raise DataAgentInputError('loaded_sources.sources must be an array')
return [source for source in sources if isinstance(source, dict)]
if paths is None:
raise DataAgentInputError('paths or loaded_sources is required')
return load_input_sources(root, paths)['sources']
def _table_headers(table: dict[str, Any]) -> list[str]:
rows = table.get('rows')
if not isinstance(rows, list) or not rows:
return []
first = rows[_table_header_index(table)]
if not isinstance(first, list):
return []
headers = []
for index, cell in enumerate(first):
text = str(cell).strip()
headers.append(text or f'column_{index + 1}')
return headers
def _table_header_index(table: dict[str, Any]) -> int:
rows = table.get('rows')
if not isinstance(rows, list):
return 0
best_index = 0
best_score = -1
for index, row in enumerate(rows[:12]):
if not isinstance(row, list):
continue
score = 0
for cell in [str(value).strip().lower() for value in row]:
if any(token in cell for token in ('query', '用户query', '预期', 'label', 'domain', '类型', '类别', '场景', '意图')):
score += 2
if any(token in cell for token in ('备注', 'case问题', '高置信', '低置信', '用户主要诉求')):
score += 1
if score > best_score:
best_score = score
best_index = index
return best_index if best_score > 0 else 0
def _profile_headers(
headers: list[str],
rows: list[Any],
field_mapping: dict[str, str],
) -> dict[str, Any]:
column_profiles: list[dict[str, Any]] = []
role_to_column: dict[str, str] = {}
explicit = {str(role): str(column) for role, column in field_mapping.items()}
for header in headers:
if not header:
continue
samples = _sample_values(header, headers, rows)
role, confidence = _detect_column_role(header, samples)
for explicit_role, explicit_column in explicit.items():
if explicit_column == header:
role = explicit_role
confidence = 1.0
if role and role not in role_to_column:
role_to_column[role] = header
column_profiles.append(
{
'name': header,
'detected_role': role or 'unknown',
'confidence': confidence,
'sample_values': samples[:5],
}
)
return {'columns': column_profiles, 'role_to_column': role_to_column}
def _detect_column_role(header: str, samples: list[str]) -> tuple[str, float]:
name = header.strip().lower()
joined = ' '.join(samples[:8]).lower()
if any(token in name for token in ('query', 'utterance', '用户query', '请求', '问题')):
return QUERY_ROLE, 0.95
if any(token in name for token in ('预期', 'expected', '正确', '人工', 'gold', 'target', 'label')):
return EXPECTED_LABEL_ROLE, 0.85
if any(token in name for token in ('pred', '预测', '模型', 'prev-domain', '当前domain', '实际domain', '命中')):
return PREDICTED_LABEL_ROLE, 0.8
if any(token in name for token in ('context', '上下文', 'prev_session', 'session')):
return CONTEXT_ROLE, 0.85
if any(token in name for token in ('备注', 'note', '原因', '反馈', 'case问题', '说明')):
return NOTES_ROLE, 0.8
if any(token in name for token in ('type', '类型', '类别', '场景', '子类')):
return CASE_TYPE_ROLE, 0.75
if any(marker in joined for marker in ('agent(', 'function', 'productagent', 'complex_task')):
return EXPECTED_LABEL_ROLE, 0.45
if name.startswith('column_') and any(len(sample) >= 10 for sample in samples):
return NOTES_ROLE, 0.35
return '', 0.0
def _questions_for_profile(profile: dict[str, Any]) -> list[str]:
role_to_column = profile.get('role_to_column') if isinstance(profile.get('role_to_column'), dict) else {}
title = profile.get('title') or profile.get('source_ref') or '表格'
questions = []
if QUERY_ROLE not in role_to_column:
questions.append(f'{title} 没有稳定识别到 query 列,请确认哪一列是用户 query。')
if EXPECTED_LABEL_ROLE not in role_to_column:
questions.append(f'{title} 没有稳定识别到预期标签列,请确认哪一列是正确 label/target。')
return questions
def _rows_to_case_evidence(
source: dict[str, Any],
table: dict[str, Any],
rows: list[Any],
role_to_column: dict[str, str],
*,
header_index: int,
max_cases: int,
) -> list[dict[str, Any]]:
if max_cases <= 0:
return []
headers = _table_headers(table)
evidence = []
for row_index, row in enumerate(rows[header_index + 1 :], start=header_index + 2):
if not isinstance(row, list):
continue
query = _cell_by_header(row, headers, role_to_column.get(QUERY_ROLE))
expected = _cell_by_header(row, headers, role_to_column.get(EXPECTED_LABEL_ROLE))
if not query:
continue
evidence.append(
{
'query': _clean_query_text(query),
'expected_label': expected,
'predicted_label': _cell_by_header(row, headers, role_to_column.get(PREDICTED_LABEL_ROLE)),
'context': _cell_by_header(row, headers, role_to_column.get(CONTEXT_ROLE)),
'case_type': _cell_by_header(row, headers, role_to_column.get(CASE_TYPE_ROLE)),
'notes': _cell_by_header(row, headers, role_to_column.get(NOTES_ROLE)),
'source_ref': f'{table.get("source_ref") or source.get("path")}:row{row_index}',
}
)
if len(evidence) >= max_cases:
break
return evidence
def _cell_by_header(row: list[Any], headers: list[str], header: str | None) -> str:
if not header:
return ''
try:
index = headers.index(header)
except ValueError:
return ''
if index >= len(row):
return ''
return str(row[index] or '').strip()
def _sample_values(header: str, headers: list[str], rows: list[Any]) -> list[str]:
samples = []
start_index = 0
for index, row in enumerate(rows[:12]):
if isinstance(row, list) and _looks_like_header_row(row, headers):
start_index = index + 1
break
for row in rows[start_index : start_index + 12]:
if not isinstance(row, list):
continue
value = _cell_by_header(row, headers, header)
if value:
samples.append(value[:120])
return samples
def _looks_like_header_row(row: list[Any], headers: list[str]) -> bool:
return [str(cell).strip() for cell in row] == headers
def _clean_query_text(text: str) -> str:
text = re.sub(r'^[QA]\s*[:]\s*', '', text.strip())
return ' / '.join(part.strip() for part in re.split(r'\s*/\s*', text) if part.strip())
def _render_table_row(headers: list[str], row: list[Any]) -> str:
pairs = []
for index, cell in enumerate(row):
text = str(cell or '').strip()
if not text:
continue
header = headers[index] if index < len(headers) else f'column_{index + 1}'
pairs.append(f'{header}: {text}')
return ' | '.join(pairs)
def _matches_focus(text: str, keywords: list[str]) -> bool:
if not keywords:
return True
return any(keyword in text for keyword in keywords)
def _clean_cell(value: Any, *, max_cell_chars: int) -> str:
if value is None:
return ''
text = ' / '.join(str(value).split())
if len(text) > max_cell_chars:
return text[: max_cell_chars - 3] + '...'
return text
def _source_summary(source: dict[str, Any]) -> str:
return (
f'{source.get("kind")} 文件,'
f'{len(source.get("paragraphs", []))} 个段落,'
f'{len(source.get("tables", []))} 个表格,'
f'{len(source.get("warnings", []))} 个警告'
)
def _dedupe(values: list[str]) -> list[str]:
seen = set()
output = []
for value in values:
if value in seen:
continue
seen.add(value)
output.append(value)
return output
def _dedupe_paths(paths: list[Path]) -> list[Path]:
seen = set()
output = []
for path in paths:
if path in seen:
continue
seen.add(path)
output.append(path)
return output
def dumps_payload(payload: dict[str, Any]) -> str:
return json.dumps(payload, ensure_ascii=False, indent=2)