The earlier report had no previous-turn rows. This run measures V0, V1, V2, and Flash on the re-extracted corpus and records that the real sample is not representative of the simulated set. Co-authored-by: Cursor <cursoragent@cursor.com>
340 lines
12 KiB
Python
340 lines
12 KiB
Python
#!/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())
|