"""Import script templates from the Sprint 2 Markdown dataset.""" from __future__ import annotations import argparse import sys from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path ROOT = Path(__file__).resolve().parents[2] sys.path.insert(0, str(ROOT)) from app.advisor_db import AgentSessionLocal from app.model.entities_advisor import ScriptTemplate from app.model.advisor_schemas import ComplianceCheckRequest from app.service.compliance_check_service import ComplianceCheckService DEFAULT_DATASET = ROOT / "docs" / "开发文档" / "29-Sprint2首批话术模板数据集.md" ACTOR_ID = "seed:template_test_data" REQUIRED_HEADERS = [ "序号", "scene", "customer_type", "title", "content", "tags", "is_active", "审核状态", "审核人", "备注", ] @dataclass(frozen=True) class SeedScriptTemplate: sequence: int scene: str customer_type: str | None title: str content: str tags: list[str] is_active: bool review_status: str reviewer: str | None note: str @property def is_approved(self) -> bool: return self.review_status == "approved" @dataclass(frozen=True) class ImportResult: total: int created: int updated: int def load_templates_from_markdown(path: Path = DEFAULT_DATASET) -> list[SeedScriptTemplate]: rows = _read_markdown_table(path) templates: list[SeedScriptTemplate] = [] for row in rows: sequence = int(row["序号"]) scene = row["scene"].strip() title = row["title"].strip() content = row["content"].strip() if not scene: raise ValueError(f"Missing scene at row {sequence}") if not title: raise ValueError(f"Missing title at row {sequence}") if not content: raise ValueError(f"Missing content at row {sequence}") templates.append( SeedScriptTemplate( sequence=sequence, scene=scene, customer_type=row["customer_type"].strip() or None, title=title, content=content, tags=_parse_tags(row["tags"]), is_active=_parse_bool(row["is_active"]), review_status=_normalize_review_status(row["审核状态"], row["备注"]), reviewer=row["审核人"].strip() or None, note=row["备注"].strip(), ) ) return templates def import_templates_from_markdown( path: Path = DEFAULT_DATASET, *, actor_id: str = ACTOR_ID, ) -> ImportResult: templates = load_templates_from_markdown(path) compliance_service = ComplianceCheckService() created = 0 updated = 0 with AgentSessionLocal() as session: for seed_template in templates: check_result = compliance_service.check_text( ComplianceCheckRequest(text=seed_template.content, scene=seed_template.scene) ) if check_result.risk_level == "BLOCK": raise ValueError(f"Template row {seed_template.sequence} failed BLOCK compliance check") existing = ( session.query(ScriptTemplate) .filter( ScriptTemplate.created_by == actor_id, ScriptTemplate.scene == seed_template.scene, ScriptTemplate.title == seed_template.title, ) .one_or_none() ) values = { "customer_type": seed_template.customer_type, "content": seed_template.content, "tags": seed_template.tags, "is_active": seed_template.is_active, "is_approved": seed_template.is_approved, "approved_by": seed_template.reviewer if seed_template.is_approved else None, "approved_at": datetime.now(timezone.utc).replace(tzinfo=None) if seed_template.is_approved else None, "updated_by": actor_id, } if existing is None: session.add( ScriptTemplate( scene=seed_template.scene, title=seed_template.title, created_by=actor_id, version=1, usage_count=0, **values, ) ) created += 1 else: if _content_changed(existing, values): existing.version += 1 for key, value in values.items(): setattr(existing, key, value) updated += 1 session.commit() return ImportResult(total=len(templates), created=created, updated=updated) def _content_changed(existing: ScriptTemplate, values: dict) -> bool: return any( [ existing.customer_type != values["customer_type"], existing.content != values["content"], existing.tags != values["tags"], existing.is_active != values["is_active"], existing.is_approved != values["is_approved"], ] ) def _read_markdown_table(path: Path) -> list[dict[str, str]]: if not path.exists(): raise FileNotFoundError(path) table_lines = [ line.strip() for line in path.read_text(encoding="utf-8").splitlines() if line.strip().startswith("|") and line.strip().endswith("|") ] if len(table_lines) < 3: raise ValueError("Markdown table is missing or empty") headers = _split_markdown_row(table_lines[0]) if headers != REQUIRED_HEADERS: raise ValueError(f"Unexpected table headers: {headers}") rows: list[dict[str, str]] = [] for line in table_lines[2:]: cells = _split_markdown_row(line) if len(cells) != len(headers): raise ValueError(f"Column count mismatch: {line}") rows.append(dict(zip(headers, cells, strict=True))) return rows def _split_markdown_row(line: str) -> list[str]: return [cell.strip() for cell in line.strip().strip("|").split("|")] def _parse_tags(value: str) -> list[str]: return [tag.strip() for tag in value.split(";") if tag.strip()] def _parse_bool(value: str) -> bool: normalized = value.strip().lower() if normalized == "true": return True if normalized == "false": return False raise ValueError(f"Expected true/false, got: {value}") def _normalize_review_status(value: str, note: str) -> str: if "已审核" in value: return "approved" if "测试数据" in value or "测试数据" in note: return "test_data" return "pending" def main() -> None: parser = argparse.ArgumentParser(description="Import Sprint 2 script templates.") parser.add_argument("--dataset", type=Path, default=DEFAULT_DATASET) args = parser.parse_args() result = import_templates_from_markdown(args.dataset) print(f"Imported script templates: total={result.total}, created={result.created}, updated={result.updated}") if __name__ == "__main__": main()