#!/usr/bin/env python3 """Build source B v2 rows with the previous turn in the same case. Offline. Writes only under the gitignored cache. Does not label gold with a model or with a keyword table. Old gold is copied by turn_id, or by a unique exact user_message when the 09-19 file has no turn_id. """ from __future__ import annotations import argparse import json import sys from datetime import datetime from pathlib import Path from typing import Any, Mapping, Sequence ROOT = Path(__file__).resolve().parents[2] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from scripts.research.jev_intent_cache import cache_dir # noqa: E402 CACHE_DIR = cache_dir() LEGACY_PATH = CACHE_DIR / "source_b.jsonl" OUTPUT_PATH = CACHE_DIR / "source_b_v2.jsonl" ANSWER_CLASSES = {"yes", "weak_yes", "no", "unsure"} CLOSED_FOCUS = {"resolved", "declined", "skipped", "superseded"} EXTRACT_SQL = """ select t.id as turn_id, t.case_id, t.user_message, t.assistant_message, t.status as turn_status, t.created_at, c.status as case_status, f.question_id, f.intent as focus_intent, f.expected_answer_schema, f.status as focus_status, f.asked_at, f.resolved_at, t.message_origin from public.agentic_rectification_turns t join public.agentic_rectification_cases c on c.id = t.case_id left join lateral ( select * from public.agentic_rectification_conversation_focuses f where f.case_id = t.case_id and f.asked_at <= t.created_at order by f.asked_at desc limit 1 ) f on true where t.user_message is not null and length(btrim(t.user_message)) > 0 order by t.case_id, t.created_at """ def _timestamp(value: Any) -> datetime | None: if value is None or value == "": return None text = str(value).strip().replace("Z", "+00:00") try: return datetime.fromisoformat(text) except ValueError: return None def _choice_copy(schema: Mapping[str, Any]) -> tuple[str, list[dict[str, str]]] | None: choice = schema.get("choice") if isinstance(schema.get("choice"), dict) else schema if not isinstance(choice, dict): return None prompt = choice.get("prompt") options = choice.get("options") if not isinstance(prompt, str) or not prompt.strip(): return None if not isinstance(options, list) or len(options) != 4: return None cleaned: list[dict[str, str]] = [] for item in options: if not isinstance(item, dict): return None key = item.get("key") label = item.get("label") answer = item.get("answer_class") if key not in {"A", "B", "C", "D"} or not isinstance(label, str) or not label.strip(): return None if answer not in ANSWER_CLASSES: return None cleaned.append({"key": str(key), "label": label.strip(), "answer_class": str(answer)}) if len({item["key"] for item in cleaned}) != 4: return None if len({item["label"] for item in cleaned}) != 4: return None return prompt.strip(), cleaned def focus_is_open(row: Mapping[str, Any]) -> bool: """Production classifies the focus that is still open when the user speaks. A focus asked earlier and already resolved, declined, skipped, or superseded is not the current question. The row stays in the extract; its layer is none. """ schema = row.get("schema") if schema is None: schema = row.get("expected_answer_schema") if not isinstance(schema, dict) or not schema: return False created = _timestamp(row.get("created_at")) asked = _timestamp(row.get("asked_at")) resolved = _timestamp(row.get("resolved_at")) if asked and created and asked > created: return False if resolved and created and resolved < created: return False status = str(row.get("focus_status") or "") if status in CLOSED_FOCUS and resolved and created and resolved < created: return False return True def focus_payload(row: Mapping[str, Any]) -> tuple[dict[str, Any] | None, str, bool]: schema = row.get("schema") if schema is None: schema = row.get("expected_answer_schema") stale = isinstance(schema, dict) and bool(schema) and not focus_is_open(row) if not focus_is_open(row) or not isinstance(schema, dict): return None, "none", stale case_status = str(row.get("case_status") or "collecting_evidence") choice = _choice_copy(schema) if choice: prompt, options = choice return { "current_question": prompt, "options": options, "case_status": case_status, "question_id": row.get("question_id"), }, "choice", False prompt = schema.get("prompt") if schema.get("collect") is True else None if isinstance(prompt, str) and prompt.strip(): return { "current_question": prompt.strip(), "options": [], "case_status": case_status, "question_id": row.get("question_id"), }, "collect", False return None, "none", False def normalize_extract_row(raw: Mapping[str, Any]) -> dict[str, Any]: focus, layer, stale = focus_payload(raw) turn_id = str(raw.get("turn_id") or raw.get("id") or "") return { "id": turn_id, "turn_id": turn_id, "source": "B", "layer": layer, "case_id": str(raw.get("case_id") or "") or None, "user_message": raw.get("user_message"), "assistant_message": raw.get("assistant_message"), "created_at": raw.get("created_at"), "case_status": raw.get("case_status"), "turn_status": raw.get("turn_status"), "message_origin": raw.get("message_origin"), "question_id": raw.get("question_id"), "focus_intent": raw.get("focus_intent"), "focus_status": raw.get("focus_status"), "focus": focus, "focus_stale": stale, } def prepare_extract(raw_rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]: return attach_previous_turns([normalize_extract_row(row) for row in raw_rows]) def load_jsonl(path: Path) -> list[dict[str, Any]]: if not path.is_file(): return [] return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] def write_jsonl(path: Path, rows: Sequence[Mapping[str, Any]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) path.write_text( "".join(json.dumps(row, ensure_ascii=False) + "\n" for row in rows), encoding="utf-8", ) def attach_previous_turns(rows: Sequence[dict[str, Any]]) -> list[dict[str, Any]]: """First turn of a case, and any row without case_id, gets previous_turn null. Rows that share a missing case_id are not chained: that would glue different people together. """ by_case: dict[str, list[dict[str, Any]]] = {} for row in rows: case_id = str(row.get("case_id") or "").strip() if not case_id: row["previous_turn"] = None row["previous_turn_source"] = "no_case" continue by_case.setdefault(case_id, []).append(row) for case_rows in by_case.values(): case_rows.sort(key=lambda item: ( str(item.get("created_at") or ""), str(item.get("turn_id") or item.get("id") or ""), )) for index, row in enumerate(case_rows): if index == 0: row["previous_turn"] = None row["previous_turn_source"] = "case_first" continue prev = case_rows[index - 1] row["previous_turn"] = { "assistant_message": "" if prev.get("assistant_message") is None else str(prev.get("assistant_message")), "user_message": "" if prev.get("user_message") is None else str(prev.get("user_message")), } row["previous_turn_source"] = "prior_turn" return list(rows) def match_gold(rows: Sequence[dict[str, Any]], legacy: Sequence[Mapping[str, Any]]) -> dict[str, int]: by_turn: dict[str, Mapping[str, Any]] = {} by_message: dict[str, list[Mapping[str, Any]]] = {} for old in legacy: turn_id = str(old.get("turn_id") or "").strip() if turn_id: by_turn[turn_id] = old message = old.get("user_message") if isinstance(message, str) and message: by_message.setdefault(message, []).append(old) counts = {"turn_id": 0, "message_exact": 0, "unlabeled": 0, "message_ambiguous": 0} for row in rows: gold = None source = "unlabeled" legacy_id = None turn_id = str(row.get("turn_id") or "").strip() if turn_id and turn_id in by_turn: gold = by_turn[turn_id].get("gold") source = "turn_id" legacy_id = by_turn[turn_id].get("id") else: hits = by_message.get(str(row.get("user_message") or "")) or [] if len(hits) == 1 and isinstance(hits[0].get("gold"), dict): gold = hits[0]["gold"] source = "message_exact" legacy_id = hits[0].get("id") elif len(hits) > 1: source = "message_ambiguous" if isinstance(gold, dict) and gold.get("intent"): row["gold"] = { "intent": gold.get("intent"), "answer_class": gold.get("answer_class"), "has_new_dated_event": gold.get("has_new_dated_event"), } row["gold_source"] = source if legacy_id: row["legacy_id"] = legacy_id else: row["gold_source"] = source row.pop("gold", None) counts[source] = counts.get(source, 0) + 1 return counts def from_legacy(legacy: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]: """The handed-over 157 rows have no case_id. previous_turn stays null.""" rows: list[dict[str, Any]] = [] for old in legacy: if not isinstance(old.get("gold"), dict) or not old["gold"].get("intent"): continue rows.append({ "id": old.get("id"), "source": "B", "layer": old.get("layer"), "user_message": old.get("user_message"), "assistant_message": None, "case_id": None, "turn_id": old.get("turn_id"), "created_at": old.get("created_at"), "focus": old.get("focus"), "focus_stale": old.get("focus_stale"), "runtime_intent": old.get("runtime_intent"), "gold": old.get("gold"), "gold_source": "legacy_file", "origin": old.get("origin"), }) return attach_previous_turns(rows) def main(argv: Sequence[str] | None = None) -> int: parser = argparse.ArgumentParser() parser.add_argument("--from-legacy", action="store_true") parser.add_argument("--raw", type=Path, default=None) parser.add_argument("--out", type=Path, default=OUTPUT_PATH) args = parser.parse_args(argv) legacy = load_jsonl(LEGACY_PATH) if args.from_legacy: rows = from_legacy(legacy) write_jsonl(args.out, rows) print(json.dumps({ "out": str(args.out), "n": len(rows), "n_with_previous": sum(1 for row in rows if row.get("previous_turn")), "gold_source": "legacy_file", }, ensure_ascii=False)) return 0 if args.raw: raw_rows = load_jsonl(args.raw) if raw_rows and any(key in raw_rows[0] for key in ("schema", "expected_answer_schema")): rows = prepare_extract(raw_rows) else: rows = attach_previous_turns([dict(row) for row in raw_rows]) counts = match_gold(rows, legacy) write_jsonl(args.out, rows) layers: dict[str, int] = {} for row in rows: layer = str(row.get("layer") or "none") layers[layer] = layers.get(layer, 0) + 1 print(json.dumps({ "out": str(args.out), "n": len(rows), "n_with_previous": sum(1 for row in rows if row.get("previous_turn")), "n_stale_focus": sum(1 for row in rows if row.get("focus_stale")), "layers": layers, "gold": counts, }, ensure_ascii=False)) return 0 print("pass --from-legacy or --raw; staging SQL is EXTRACT_SQL in this file", file=sys.stderr) return 2 if __name__ == "__main__": raise SystemExit(main())