From 2eb05392c6158a44d56f73b43a1789983d521540 Mon Sep 17 00:00:00 2001 From: jinliyl <6469360+jinliyl@users.noreply.github.com> Date: Thu, 16 Jul 2026 20:32:28 +0800 Subject: [PATCH] chore(benchmark): remove longmemeval final answer review file (#366) * feat(benchmark): add final answer review step for evaluation - Introduce FinalAnswerReviewStep to handle answer validation - Add final_answer_review.jsonl dataset with 24 evaluation cases - Include detailed reasoning and golden check results for each case - Support various question types including temporal reasoning and preferences - Implement time consistency checks for session references - Add comprehensive test coverage for different evaluation scenarios * chore(benchmark): remove longmemeval final answer review file - Removed final_answer_review.jsonl containing 23 evaluation records - Deleted question_id mappings with detailed reasoning for golden answers - Removed answer correctness assessments and session time validation checks - Cleaned up benchmark dataset used for memory evaluation testing - Eliminated JSONL format evaluation results for temporal reasoning tasks - Removed references to various session IDs and time-based validations * config(default): disable shell step configuration by commenting out - Commented out the shell step configuration in default.yaml - Disabled asynchronous shell command execution capability - Removed shell step from available backend operations - Preserved traverse backend configuration unchanged * refactor(tests): remove unused shell job test from config parser tests - Removed test_default_config_registers_shell_job function that was no longer needed - Kept existing test for frontmatter chunk metadata configuration - Cleaned up test suite by removing obsolete test case --- .../longmemeval/run_final_answer_review.py | 469 ++++++++++++++++++ reme/config/default.yaml | 34 +- reme/config/jinli_lme.yaml | 23 + reme/steps/benchmark/lme/__init__.py | 2 + .../benchmark/lme/final_answer_review.py | 280 +++++++++++ .../benchmark/lme/final_answer_review.yaml | 53 ++ reme/steps/benchmark/lme/golden_check.py | 31 +- reme/steps/benchmark/lme/golden_check.yaml | 16 +- reme/steps/benchmark/lme/session_review.yaml | 3 +- tests/unit/test_config_parser.py | 11 - tests/unit/test_lme_final_answer_review.py | 425 ++++++++++++++++ 11 files changed, 1285 insertions(+), 62 deletions(-) create mode 100644 benchmark/longmemeval/run_final_answer_review.py create mode 100644 reme/steps/benchmark/lme/final_answer_review.py create mode 100644 reme/steps/benchmark/lme/final_answer_review.yaml create mode 100644 tests/unit/test_lme_final_answer_review.py diff --git a/benchmark/longmemeval/run_final_answer_review.py b/benchmark/longmemeval/run_final_answer_review.py new file mode 100644 index 00000000..2c8491ad --- /dev/null +++ b/benchmark/longmemeval/run_final_answer_review.py @@ -0,0 +1,469 @@ +#!/usr/bin/env python3 +"""Review every LongMemEval golden answer with the configured Claude Code job. + +Every numeric ``datasets/longmemeval/`` workspace is processed sequentially. +The reference JSONL files are merged by ``question_id`` and supplied only when +they contain an alternative answer for that sample: + + reme start config=jinli_lme job=final_answer_review + +The job returns a plain four-field JSON object with ``reason``, +``golden_answer_correct``, ``answer``, and ``is_session_time_wrong``. After +every new success, this driver atomically rewrites the complete accumulated +output JSONL so an interrupted run can safely resume. + +Examples: + python benchmark/longmemeval/run_final_answer_review.py + python benchmark/longmemeval/run_final_answer_review.py --exclude-reference-question-ids + python benchmark/longmemeval/run_final_answer_review.py --only-reference-question-ids --rerun-selected + python benchmark/longmemeval/run_final_answer_review.py --concurrency 2 --submit-interval-seconds 6 + python benchmark/longmemeval/run_final_answer_review.py --question-id e47becba + python benchmark/longmemeval/run_final_answer_review.py --reference path/to/results.jsonl + python benchmark/longmemeval/run_final_answer_review.py --limit 3 + python benchmark/longmemeval/run_final_answer_review.py --no-resume + python benchmark/longmemeval/run_final_answer_review.py --dry-run +""" + +import argparse +import concurrent.futures +import json +import os +import subprocess +import sys +import tempfile +import time +from pathlib import Path +from typing import Any + +REPO = Path(__file__).resolve().parents[2] +DATA = REPO / "datasets" / "longmemeval" +DEFAULT_REFERENCES = ( + REPO / "benchmark" / "longmemeval" / "golden_check_list_false.jsonl", + REPO / "benchmark" / "longmemeval" / "merge_confirm_jinli_false.jsonl", +) +DEFAULT_OUTPUT = REPO / "benchmark" / "longmemeval" / "final_answer_review.jsonl" +DEFAULT_LOG_DIR = REPO / "logs" / "final_answer_review" +REFERENCE_PATHS_ENV = "LME_FINAL_ANSWER_REFERENCE_PATHS" +MAX_CONCURRENCY = 3 +MIN_SUBMIT_INTERVAL_SECONDS = 5.0 +DEFAULT_SUBMIT_INTERVAL_SECONDS = 5.1 + + +def parse_args() -> argparse.Namespace: + """Parse command-line arguments.""" + parser = argparse.ArgumentParser( + description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + parser.add_argument( + "--question-id", + dest="question_ids", + action="append", + help="process only this dataset question ID; repeat for multiple IDs (default: all)", + ) + reference_selection = parser.add_mutually_exclusive_group() + reference_selection.add_argument( + "--exclude-reference-question-ids", + action="store_true", + help="skip question IDs found in the selected reference-answer JSONL files", + ) + reference_selection.add_argument( + "--only-reference-question-ids", + action="store_true", + help="process only question IDs found in the selected reference-answer JSONL files", + ) + parser.add_argument( + "--reference", + dest="references", + action="append", + type=Path, + help="reference-answer JSONL; repeat for multiple files (default: built-in disputed results)", + ) + parser.add_argument( + "--output", + type=Path, + default=DEFAULT_OUTPUT, + help=f"output JSONL (default: {DEFAULT_OUTPUT})", + ) + parser.add_argument( + "--log-dir", + type=Path, + default=DEFAULT_LOG_DIR, + help="directory for per-question logs", + ) + parser.add_argument( + "--concurrency", + type=int, + default=MAX_CONCURRENCY, + help=f"maximum concurrent jobs, from 1 to {MAX_CONCURRENCY} (default: {MAX_CONCURRENCY})", + ) + parser.add_argument( + "--submit-interval-seconds", + type=float, + default=DEFAULT_SUBMIT_INTERVAL_SECONDS, + help=f"minimum time between job submissions; must be > {MIN_SUBMIT_INTERVAL_SECONDS:g} " + f"(default: {DEFAULT_SUBMIT_INTERVAL_SECONDS:g})", + ) + parser.add_argument( + "--limit", + type=int, + default=0, + help="process only the first N pending questions (0 = all)", + ) + resume_mode = parser.add_mutually_exclusive_group() + resume_mode.add_argument( + "--no-resume", + action="store_true", + help="ignore existing output and rerun every selected question", + ) + resume_mode.add_argument( + "--rerun-selected", + action="store_true", + help="rerun every selected question while preserving existing results until replacements finish", + ) + parser.add_argument( + "--dry-run", + action="store_true", + help="show the selected cases without invoking ReMe", + ) + return parser.parse_args() + + +def _read_jsonl(path: Path) -> list[dict[str, Any]]: + """Read a JSONL file and reject malformed or non-object rows.""" + rows: list[dict[str, Any]] = [] + try: + with path.open(encoding="utf-8") as file: + for line_number, line in enumerate(file, start=1): + if not line.strip(): + continue + try: + row = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError(f"Invalid JSON at {path}:{line_number}") from exc + if not isinstance(row, dict): + raise ValueError(f"Expected a JSON object at {path}:{line_number}") + rows.append(row) + except OSError as exc: + raise FileNotFoundError(f"Cannot read JSONL file: {path}") from exc + return rows + + +def merge_references(paths: list[Path]) -> dict[str, list[dict[str, Any]]]: + """Merge reference rows by question ID, preserving file and row order.""" + merged: dict[str, list[dict[str, Any]]] = {} + seen_sources: set[tuple[str, str]] = set() + for path in paths: + for row in _read_jsonl(path): + question_id = str(row.get("question_id") or "").strip() + if not question_id: + raise ValueError(f"Reference row in {path} has no question_id") + source_key = (question_id, str(path.resolve())) + if source_key in seen_sources: + raise ValueError(f"Duplicate question_id={question_id!r} within {path}") + seen_sources.add(source_key) + merged.setdefault(question_id, []).append({"source": path.name, **row}) + if not merged: + raise ValueError("No reference answers found") + return merged + + +def workspace_map() -> dict[str, Path]: + """Map every dataset question ID to its numeric sample workspace.""" + mapping: dict[str, Path] = {} + for workspace in sorted( + (path for path in DATA.iterdir() if path.is_dir() and path.name.isdigit()), + key=lambda p: int(p.name), + ): + query_path = workspace / "query.json" + if not query_path.is_file(): + continue + try: + with query_path.open(encoding="utf-8") as file: + query = json.load(file) + except (OSError, json.JSONDecodeError) as exc: + raise ValueError(f"Cannot parse {query_path}") from exc + if not isinstance(query, dict): + raise ValueError(f"Expected a JSON object in {query_path}") + question_id = str(query.get("question_id") or "").strip() + if not question_id: + raise ValueError(f"Missing question_id in {query_path}") + if question_id in mapping: + raise ValueError( + f"Duplicate dataset question_id={question_id!r}: {mapping[question_id]} and {workspace}", + ) + mapping[question_id] = workspace + return mapping + + +def select_question_ids( + mapping: dict[str, Path], + requested: list[str] | None, + excluded: set[str] | None = None, +) -> list[str]: + """Return all dataset IDs or validate an explicitly requested subset.""" + excluded = excluded or set() + if not requested: + return [question_id for question_id in mapping if question_id not in excluded] + selected: list[str] = [] + seen: set[str] = set() + for raw_question_id in requested: + question_id = raw_question_id.strip() + if not question_id: + raise ValueError("--question-id must not be empty") + if question_id in seen: + raise ValueError(f"Duplicate --question-id: {question_id}") + if question_id not in mapping: + raise ValueError(f"No dataset workspace for question ID: {question_id}") + if question_id not in excluded: + selected.append(question_id) + seen.add(question_id) + return selected + + +def _validate_result(value: Any, *, source: str) -> dict[str, Any]: + """Validate the final four-field answer contract.""" + expected_keys = {"reason", "golden_answer_correct", "answer", "is_session_time_wrong"} + if not isinstance(value, dict) or set(value) != expected_keys: + raise ValueError( + f"{source} must contain exactly 'reason', 'golden_answer_correct', 'answer', " + "and 'is_session_time_wrong'", + ) + if not isinstance(value["reason"], str) or not value["reason"].strip(): + raise ValueError(f"{source} has an invalid reason") + if not isinstance(value["golden_answer_correct"], bool): + raise ValueError(f"{source} has an invalid golden_answer_correct") + if not isinstance(value["answer"], str): + raise ValueError(f"{source} has an invalid answer") + answer = value["answer"].strip() + if value["golden_answer_correct"] and answer: + raise ValueError(f"{source} answer must be empty when golden_answer_correct is true") + if not value["golden_answer_correct"] and not answer: + raise ValueError(f"{source} answer must be non-empty when golden_answer_correct is false") + if not isinstance(value["is_session_time_wrong"], bool): + raise ValueError(f"{source} has an invalid is_session_time_wrong") + return { + "reason": value["reason"].strip(), + "golden_answer_correct": value["golden_answer_correct"], + "answer": answer, + "is_session_time_wrong": False, + } + + +def load_existing(path: Path) -> dict[str, dict[str, Any]]: + """Load resumable output, rejecting duplicate or malformed rows.""" + if not path.exists(): + return {} + results: dict[str, dict[str, Any]] = {} + for row in _read_jsonl(path): + question_id = str(row.get("question_id") or "").strip() + if not question_id: + raise ValueError(f"Existing output row in {path} has no question_id") + if question_id in results: + raise ValueError( + f"Duplicate question_id={question_id!r} in existing output {path}", + ) + results[question_id] = _validate_result( + {key: value for key, value in row.items() if key != "question_id"}, + source=f"existing result for {question_id}", + ) + return results + + +def atomic_write_results( + path: Path, + order: list[str], + results: dict[str, dict[str, Any]], +) -> None: + """Atomically rewrite all accumulated rows in stable merged-input order.""" + path.parent.mkdir(parents=True, exist_ok=True) + temp_path: Path | None = None + try: + with tempfile.NamedTemporaryFile( + "w", + encoding="utf-8", + dir=path.parent, + prefix=f".{path.name}.", + delete=False, + ) as file: + temp_path = Path(file.name) + for question_id in order: + if question_id not in results: + continue + row = {"question_id": question_id, **results[question_id]} + file.write( + json.dumps(row, ensure_ascii=False, separators=(",", ":")) + "\n", + ) + file.flush() + os.fsync(file.fileno()) + os.replace(temp_path, path) + finally: + if temp_path is not None and temp_path.exists(): + temp_path.unlink() + + +def run_one( + question_id: str, + workspace: Path, + log_dir: Path, + reference_paths: list[Path], +) -> dict[str, Any]: + """Run the configured one-shot job and validate its stdout JSON.""" + env = dict(os.environ, LME_WORKSPACE_DIR=str(workspace.relative_to(REPO))) + env[REFERENCE_PATHS_ENV] = json.dumps( + [str(path.resolve()) for path in reference_paths], + ensure_ascii=False, + ) + completed = subprocess.run( + [ + sys.executable, + "-c", + "from reme.reme import main; main()", + "start", + "config=jinli_lme", + "job=final_answer_review", + ], + cwd=REPO, + env=env, + text=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + check=False, + ) + log_dir.mkdir(parents=True, exist_ok=True) + log_path = log_dir / f"{question_id}.log" + log_text = ( + f"workspace={workspace}\nreturncode={completed.returncode}\n\n" + f"[stdout]\n{completed.stdout}\n[stderr]\n{completed.stderr}" + ) + log_path.write_text( + log_text, + encoding="utf-8", + ) + if completed.returncode != 0: + raise RuntimeError( + f"Job failed for {question_id} with rc={completed.returncode}; see {log_path}", + ) + try: + value = json.loads(completed.stdout.strip()) + except json.JSONDecodeError as exc: + raise ValueError( + f"Job stdout is not JSON for {question_id}; see {log_path}", + ) from exc + return _validate_result(value, source=f"job result for {question_id}") + + +def main() -> int: + """Review and checkpoint the selected dataset cases sequentially.""" + args = parse_args() + if args.limit < 0: + raise ValueError("--limit must be >= 0") + if not 1 <= args.concurrency <= MAX_CONCURRENCY: + raise ValueError(f"--concurrency must be between 1 and {MAX_CONCURRENCY}") + if args.submit_interval_seconds <= MIN_SUBMIT_INTERVAL_SECONDS: + raise ValueError( + f"--submit-interval-seconds must be > {MIN_SUBMIT_INTERVAL_SECONDS:g}", + ) + + reference_paths = [path.resolve() for path in (args.references or DEFAULT_REFERENCES)] + mapping = workspace_map() + references = merge_references(reference_paths) + missing = [question_id for question_id in references if question_id not in mapping] + if missing: + raise ValueError(f"No dataset workspace for question IDs: {', '.join(missing)}") + + full_order = list(mapping) + excluded = set(references) if args.exclude_reference_question_ids else set() + order = select_question_ids(mapping, args.question_ids, excluded) + if args.only_reference_question_ids: + order = [question_id for question_id in order if question_id in references] + results = {} if args.no_resume else load_existing(args.output.resolve()) + pending = ( + list(order) if args.rerun_selected else [question_id for question_id in order if question_id not in results] + ) + if args.limit: + pending = pending[: args.limit] + + no_reference = sum(question_id not in references for question_id in order) + one_reference = sum(len(references.get(question_id, [])) == 1 for question_id in order) + multiple_references = sum(len(references.get(question_id, [])) > 1 for question_id in order) + print( + f"total={len(order)} no_reference={no_reference} one_reference={one_reference} " + f"multiple_references={multiple_references} " + f"excluded={len(excluded)} " + f"only_reference_questions={args.only_reference_question_ids} " + f"concurrency={args.concurrency} submit_interval={args.submit_interval_seconds:g}s " + f"existing={len(results)} pending={len(pending)} output={args.output.resolve()}", + flush=True, + ) + + if args.dry_run: + for question_id in pending: + print( + f"[would-run] question_id={question_id} workspace={mapping[question_id].name} " + f"references={len(references.get(question_id, []))}", + ) + return 0 + + executor = concurrent.futures.ThreadPoolExecutor(max_workers=args.concurrency) + active: dict[concurrent.futures.Future[dict[str, Any]], tuple[int, str]] = {} + next_position = 0 + saved_count = 0 + next_submit_at = 0.0 + try: + while next_position < len(pending) or active: + can_submit = next_position < len(pending) and len(active) < args.concurrency + if can_submit and time.monotonic() >= next_submit_at: + question_id = pending[next_position] + position = next_position + 1 + workspace = mapping[question_id] + print( + f"[submit {position}/{len(pending)}] question_id={question_id} " + f"workspace={workspace.name} references={len(references.get(question_id, []))}", + flush=True, + ) + future = executor.submit( + run_one, + question_id, + workspace, + args.log_dir.resolve(), + reference_paths, + ) + active[future] = (position, question_id) + next_position += 1 + next_submit_at = time.monotonic() + args.submit_interval_seconds + continue + + if not active: + time.sleep(max(0.0, next_submit_at - time.monotonic())) + continue + + timeout = None + if can_submit: + timeout = max(0.0, next_submit_at - time.monotonic()) + done, _ = concurrent.futures.wait( + active, + timeout=timeout, + return_when=concurrent.futures.FIRST_COMPLETED, + ) + for future in done: + position, question_id = active.pop(future) + results[question_id] = future.result() + atomic_write_results(args.output.resolve(), full_order, results) + saved_count += 1 + print( + f"[saved {saved_count}/{len(pending)}] submitted_position={position} " f"question_id={question_id}", + flush=True, + ) + finally: + executor.shutdown(wait=True, cancel_futures=True) + + print( + f"ALL FINISHED total_saved={sum(question_id in results for question_id in order)}", + flush=True, + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/reme/config/default.yaml b/reme/config/default.yaml index 843981f0..c810758b 100644 --- a/reme/config/default.yaml +++ b/reme/config/default.yaml @@ -237,23 +237,23 @@ jobs: steps: - backend: help_step - shell: - backend: base - description: "execute a shell command asynchronously in the workspace" - parameters: - type: object - properties: - cmd: - type: string - description: "shell command to execute" - shell_timeout: - type: number - description: "maximum execution time in seconds" - default: 86400 - required: - - cmd - steps: - - backend: shell_step +# shell: +# backend: base +# description: "execute a shell command asynchronously in the workspace" +# parameters: +# type: object +# properties: +# cmd: +# type: string +# description: "shell command to execute" +# shell_timeout: +# type: number +# description: "maximum execution time in seconds" +# default: 86400 +# required: +# - cmd +# steps: +# - backend: shell_step traverse: backend: base diff --git a/reme/config/jinli_lme.yaml b/reme/config/jinli_lme.yaml index 5814e8b1..df7906aa 100644 --- a/reme/config/jinli_lme.yaml +++ b/reme/config/jinli_lme.yaml @@ -199,6 +199,21 @@ jobs: - backend: lme_golden_check_step agent_wrapper: lme_judge + final_answer_review: + backend: base + description: "Review one LongMemEval golden answer from all sessions available by question_date." + parameters: + type: object + properties: { } + steps: + - backend: lme_final_answer_review_step + agent_wrapper: lme_final_answer_review + reference_paths: + - benchmark/longmemeval/golden_check_list_false.jsonl + - benchmark/longmemeval/merge_confirm_jinli_false.jsonl + retry_initial_seconds: 5 + retry_max_seconds: 300 + components: tokenizer: default: @@ -335,6 +350,14 @@ components: base_url: ${CLAUDE_CODE_BASE_URL:-https://dashscope.aliyuncs.com/apps/anthropic} permission_mode: bypassPermissions + lme_final_answer_review: + backend: claude_code + model: ${CLAUDE_CODE_MODEL_NAME:-claude-opus-4-8} + api_key: ${CLAUDE_CODE_API_KEY:-} + base_url: ${CLAUDE_CODE_BASE_URL:-} + cwd: session + permission_mode: bypassPermissions + lme_review: backend: agentscope as_llm: plus diff --git a/reme/steps/benchmark/lme/__init__.py b/reme/steps/benchmark/lme/__init__.py index c9c5b4a3..aba19f47 100644 --- a/reme/steps/benchmark/lme/__init__.py +++ b/reme/steps/benchmark/lme/__init__.py @@ -4,12 +4,14 @@ from .agentic_answer import LmeAgenticAnswerStep from .auto_memory import LmeAutoMemoryStep from .context_answer import ContextAnswerStep from .extract_session import LmeExtractSessionStep +from .final_answer_review import FinalAnswerReviewStep from .golden_check import GoldenCheckStep from .lme_llm_judge import LmeLlmJudgeStep from .session_review import SessionReviewStep __all__ = [ "ContextAnswerStep", + "FinalAnswerReviewStep", "GoldenCheckStep", "LmeAgenticAnswerStep", "LmeAutoMemoryStep", diff --git a/reme/steps/benchmark/lme/final_answer_review.py b/reme/steps/benchmark/lme/final_answer_review.py new file mode 100644 index 00000000..5fb86289 --- /dev/null +++ b/reme/steps/benchmark/lme/final_answer_review.py @@ -0,0 +1,280 @@ +"""Produce a final, evidence-backed answer for a LongMemEval case. + +The step puts the complete query, golden-answer object, and any available +disputed reference answers directly into the prompt. Raw session content stays out of the model +context: Claude Code starts in the sample's ``session`` directory and uses its +normal file tools to inspect whichever sessions it needs. Session timestamps +are scanned only to identify evidence that did not exist at question time; +``answer_session_ids`` are not evaluated. + +Claude Code is intentionally used without an output schema. Its ordinary text +reply may contain narration but must include exactly one fenced ``json`` block +whose object contains ``reason``, ``golden_answer_correct``, ``answer``, and +``is_session_time_wrong``. API errors and invalid replies are retried with +capped exponential backoff. +""" + +import asyncio +import json +import os +import re +from datetime import datetime +from pathlib import Path +from typing import Any + +from ....components import R +from ...base_step import BaseStep + +DEFAULT_REFERENCE_PATHS = ( + "benchmark/longmemeval/golden_check_list_false.jsonl", + "benchmark/longmemeval/merge_confirm_jinli_false.jsonl", +) +REFERENCE_PATHS_ENV = "LME_FINAL_ANSWER_REFERENCE_PATHS" +RETRY_INITIAL_SECONDS = 5.0 +RETRY_MAX_SECONDS = 300.0 +_LME_DATETIME_RE = re.compile(r"(\d{4})/(\d{2})/(\d{2}).*?(\d{2}):(\d{2})") +_FENCED_JSON_RE = re.compile(r"```json\s*(.*?)\s*```", re.IGNORECASE | re.DOTALL) + + +@R.register("lme_final_answer_review_step") +class FinalAnswerReviewStep(BaseStep): + """Ask a Claude Code agent to review one golden answer.""" + + @staticmethod + def _load_json(path: Path) -> dict[str, Any]: + try: + with path.open(encoding="utf-8") as file: + value = json.load(file) + except OSError as exc: + raise FileNotFoundError(f"Cannot read LongMemEval file: {path}") from exc + except json.JSONDecodeError as exc: + raise ValueError(f"Invalid JSON in LongMemEval file: {path}") from exc + if not isinstance(value, dict): + raise ValueError(f"Expected a JSON object in {path}") + return value + + @staticmethod + def _parse_datetime(raw_date: Any, *, source: str) -> datetime: + text = str(raw_date or "").strip() + match = _LME_DATETIME_RE.search(text) + if match is None: + raise ValueError(f"Invalid LongMemEval datetime in {source}: {text!r}") + try: + return datetime(*(int(part) for part in match.groups())) + except ValueError as exc: + raise ValueError( + f"Invalid LongMemEval datetime in {source}: {text!r}", + ) from exc + + def _resolve_reference_path(self, raw_path: str) -> Path: + path = Path(raw_path).expanduser() + if path.is_absolute(): + return path + + # The configured defaults are repository-relative. Tests and custom + # jobs may instead provide workspace-relative fixture paths. + repository_path = Path.cwd() / path + if repository_path.is_file(): + return repository_path + return self.workspace_path / path + + def _load_references(self, question_id: str) -> list[dict[str, Any]]: + raw_paths: Any + serialized_paths = os.environ.get(REFERENCE_PATHS_ENV) + if serialized_paths: + try: + raw_paths = json.loads(serialized_paths) + except json.JSONDecodeError as exc: + raise ValueError(f"{REFERENCE_PATHS_ENV} must be a JSON array of paths") from exc + else: + raw_paths = self.kwargs.get("reference_paths") or DEFAULT_REFERENCE_PATHS + if isinstance(raw_paths, str): + raw_paths = [raw_paths] + if not isinstance(raw_paths, (list, tuple)) or not raw_paths: + raise ValueError("reference_paths must contain at least one JSONL path") + + references: list[dict[str, Any]] = [] + for raw_path in raw_paths: + path = self._resolve_reference_path(str(raw_path)) + try: + with path.open(encoding="utf-8") as file: + for line_number, line in enumerate(file, start=1): + if not line.strip(): + continue + try: + item = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError( + f"Invalid JSONL at {path}:{line_number}", + ) from exc + if not isinstance(item, dict): + raise ValueError( + f"Expected a JSON object at {path}:{line_number}", + ) + if str(item.get("question_id") or "") == question_id: + references.append({"source": path.name, **item}) + except OSError as exc: + raise FileNotFoundError( + f"Cannot read reference-answer file: {path}", + ) from exc + + return references + + def _inspect_session_times(self, question_dt: datetime) -> tuple[int, list[dict[str, str]]]: + """Return the session count and timestamp-only metadata for future sessions.""" + resource_dir = self.app_context.app_config.resource_dir if self.app_context is not None else "session" + session_dir = self.workspace_path / resource_dir + if not session_dir.is_dir(): + raise FileNotFoundError(f"Session directory not found: {session_dir}") + + session_paths = sorted(session_dir.glob("*.json")) + future_sessions: list[dict[str, str]] = [] + for path in session_paths: + session = self._load_json(path) + session_id = str(session.get("haystack_session_id") or path.stem) + session_date = str(session.get("haystack_date") or "").strip() + session_dt = self._parse_datetime( + session_date, + source=f"{path}:haystack_date", + ) + if session_dt > question_dt: + future_sessions.append( + { + "session_id": session_id, + "session_date": session_date, + "session_file": path.name, + }, + ) + return len(session_paths), future_sessions + + @staticmethod + def _parse_reply(raw_reply: Any) -> dict[str, Any]: + if not isinstance(raw_reply, str) or not raw_reply.strip(): + raise ValueError("Agent returned an empty reply") + json_blocks = _FENCED_JSON_RE.findall(raw_reply) + if len(json_blocks) != 1: + raise ValueError("Agent reply must contain exactly one fenced ```json``` block") + try: + value = json.loads(json_blocks[0].strip()) + except json.JSONDecodeError as exc: + raise ValueError("Agent's fenced json block is not valid JSON") from exc + if not isinstance(value, dict): + raise ValueError("Agent reply must be a JSON object") + if set(value) != {"reason", "golden_answer_correct", "answer", "is_session_time_wrong"}: + raise ValueError( + "Agent reply must contain exactly 'reason', 'golden_answer_correct', 'answer', " + "and 'is_session_time_wrong'", + ) + answer = value["answer"] + reason = value["reason"] + golden_answer_correct = value["golden_answer_correct"] + is_session_time_wrong = value["is_session_time_wrong"] + if not isinstance(reason, str) or not reason.strip(): + raise ValueError("Agent reply 'reason' must be a non-empty string") + if "answer_session_ids" in reason.casefold(): + raise ValueError("Agent reply 'reason' must not evaluate answer_session_ids") + if not isinstance(golden_answer_correct, bool): + raise ValueError("Agent reply 'golden_answer_correct' must be a boolean") + if not isinstance(answer, str): + raise ValueError("Agent reply 'answer' must be a string") + answer = answer.strip() + if golden_answer_correct and answer: + raise ValueError("Agent reply 'answer' must be empty when golden_answer_correct is true") + if not golden_answer_correct and not answer: + raise ValueError("Agent reply 'answer' must be non-empty when golden_answer_correct is false") + if not isinstance(is_session_time_wrong, bool): + raise ValueError("Agent reply 'is_session_time_wrong' must be a boolean") + if is_session_time_wrong: + raise ValueError("Agent reply 'is_session_time_wrong' is deprecated and must be false") + return { + "reason": reason.strip(), + "golden_answer_correct": golden_answer_correct, + "answer": answer, + "is_session_time_wrong": is_session_time_wrong, + } + + async def execute(self): + assert self.context is not None + if self.agent_wrapper is None: + raise ValueError("lme_final_answer_review_step requires agent_wrapper") + + query = self._load_json(self.workspace_path / "query.json") + golden = self._load_json(self.workspace_path / "answer.json") + question_id = str(query.get("question_id") or "").strip() + if not question_id: + raise ValueError("query.json requires a non-empty 'question_id'") + question_dt = self._parse_datetime( + query.get("question_date"), + source="query.json:question_date", + ) + references = self._load_references(question_id) + num_sessions, future_sessions = self._inspect_session_times(question_dt) + + payload = { + "query": query, + "answer_json": golden, + "reference_answers": references, + "session_time_check": { + "sessions_after_question_date": future_sessions, + }, + } + user_prompt = self.prompt_format( + "user_message", + question_id=question_id, + question_date=str(query.get("question_date") or ""), + num_sessions=num_sessions, + num_future_sessions=len(future_sessions), + num_references=len(references), + payload_json=json.dumps(payload, ensure_ascii=False, indent=2), + ) + + retry_initial_seconds = float( + self.kwargs.get("retry_initial_seconds", RETRY_INITIAL_SECONDS), + ) + retry_max_seconds = float( + self.kwargs.get("retry_max_seconds", RETRY_MAX_SECONDS), + ) + if retry_initial_seconds <= 0: + retry_initial_seconds = RETRY_INITIAL_SECONDS + retry_max_seconds = max(retry_max_seconds, retry_initial_seconds) + + attempt = 1 + sleep_seconds = retry_initial_seconds + while True: + try: + # Deliberately do not pass output_schema: this case evaluates an + # ordinary Claude Code response and validates it afterward. + result = await self.agent_wrapper.reply( + user_prompt, + system_prompt=self.get_prompt("system_prompt"), + ) + final_answer = self._parse_reply(result.get("result")) + if attempt > 1: + self.logger.info( + f"[{self.name}] recovered after {attempt} attempts", + ) + break + except Exception as exc: # noqa: BLE001 - agent/API/format failures share the retry contract + delay = min(sleep_seconds, retry_max_seconds) + self.logger.warning( + f"[{self.name}] attempt {attempt} failed for {question_id}: {exc}; retrying in {delay:.1f}s", + ) + await asyncio.sleep(delay) + sleep_seconds = min(sleep_seconds * 2, retry_max_seconds) + attempt += 1 + + self.context.response.success = True + self.context.response.answer = json.dumps(final_answer, ensure_ascii=False) + self.context.response.metadata.update( + { + "question_id": question_id, + "num_sessions": num_sessions, + "num_future_sessions": len(future_sessions), + "future_sessions": future_sessions, + "num_reference_answers": len(references), + "is_session_time_wrong": False, + "attempts": attempt, + "agent_session_id": result.get("session_id"), + }, + ) + return self.context.response diff --git a/reme/steps/benchmark/lme/final_answer_review.yaml b/reme/steps/benchmark/lme/final_answer_review.yaml new file mode 100644 index 00000000..613d37df --- /dev/null +++ b/reme/steps/benchmark/lme/final_answer_review.yaml @@ -0,0 +1,53 @@ +system_prompt: | + 你是 LongMemEval 答案的最终审核员。完整的 query.json、answer.json,以及零个或多个可能正确、 + 也可能错误的参考答案已经放在用户消息的 input JSON 中,不需要去其他目录寻找这些输入。没有参考 + 答案时,应直接根据原始 session 独立审核 answer.json,不能因为缺少争议记录就假定 golden 答案正确。 + + 你的当前工作目录就是该问题的 session 目录。目录中的每个 JSON 文件都是一个完整原始聊天 session。 + 原始 session 内容没有预先放进上下文;请主动使用 Read、Glob、Grep、Bash 等工具在当前目录自由检索, + 并阅读所有与问题可能相关的 session。不要修改或删除这些文件。 + + 你的任务是独立判断最合理的答案。answer.json 和 reference_answers 都只是待核对的线索,不是事实, + 不能因为多个参考答案一致就直接采纳。必须综合全部聊天记录,仔细区分用户与 assistant 的陈述,处理 + 时间、更新、冲突、计数、偏好和指代关系。 + + 检索时必须始终检查每个文件中的 haystack_date:发生在 question_date 之后的 session 属于未来 + 信息,绝对不能用其聊天内容推导正确答案或判断 answer.json 正确。即使未来 session 给出了非常直接、 + 看似正确或与参考答案一致的信息,也必须忽略其内容,避免时间穿越。必须先仅根据 question_date 当时 + 已经存在的 session 独立得出正确答案,再与 answer.json 比较;合法证据不足时,正确答案为 unknown。 + + `answer_session_ids` 不属于本次审核对象。不要检查其是否完整、相关、存在或晚于 question_date,也 + 不得因其包含未来、无关或错误的 session ID 而把 golden answer 判错。`golden_answer_correct` 只由 + `answer.json` 中 `answer` 的内容是否完整、正确决定。 + + input JSON 中的 session_time_check 只用于指出哪些 session 内容晚于 question_date、不能作为答题 + 证据;它不用于检查 `answer_session_ids`。reason 中不需要评价 `answer_session_ids`。 + + 你可以在最终回复中补充必要的分析文字,但必须包含且只能包含一个 ```json 代码块。程序只解析这个 + 代码块;没有代码块、存在多个 json 代码块或块内 JSON 无效都会触发重试。代码块内必须是一个对象, + 且只能包含四个字段: + - reason:中文详细推理。说明如何处理不同线索和参考答案,尽量逐条引用有证据作用的 session id、 + session 时间与具体事实,使后续人工 reviewer 可以复核。 + - golden_answer_correct:JSON boolean。仅根据 question_date 之前(含同一时刻)的 session 判断 + answer.json 中的 answer 是否完整且正确;不要考虑 answer_session_ids。 + - answer:仅当 golden_answer_correct 为 false 时,填写合法证据支持的正确答案(证据不足填 + unknown);为 true 时必须填空字符串。 + - is_session_time_wrong:为兼容现有输出结构保留的弃用字段,始终填 false。 + + 不要在 reason 或任何字段中评价 answer_session_ids。 + + 输出格式示例仅用于说明 JSON 外形,不是内容 few-shot: + ```json + {"reason":"详细推理与 session 证据","golden_answer_correct":false,"answer":"修正答案","is_session_time_wrong":false} + ``` + +user_message: | + 请审核 question_id={question_id}。 + Question date: {question_date} + Session files in current working directory: {num_sessions} + Sessions after question_date: {num_future_sessions} + Reference answer count: {num_references} + + 以下 input JSON 包含完整 query.json、answer.json、参考答案和 session 时间检查结果。请先读完,再使用当前 + session 目录中的原始文件查找证据,独立推理后严格按 system prompt 要求输出带 ```json 代码块的结果: + {payload_json} diff --git a/reme/steps/benchmark/lme/golden_check.py b/reme/steps/benchmark/lme/golden_check.py index 2a06d728..5e8e4720 100644 --- a/reme/steps/benchmark/lme/golden_check.py +++ b/reme/steps/benchmark/lme/golden_check.py @@ -3,8 +3,8 @@ Consumes ``session_review.json`` produced by ``lme_session_review_step`` and hands its extracted session information to an agent that is equipped with the ``python_execute`` tool. The agent uses ``python_execute`` only as a scratchpad -for the hard reasoning (checking the golden answer and cross-checking the -filtered ``answer_session_ids``); the final verdict is not the +for checking the golden answer; ``answer_session_ids`` are outside the audit +scope. The final verdict is not the raw Python stdout but a *structured* object extracted from the whole conversation via ``output_schema``. Sessions dated after ``question_date`` are filtered upstream by ``lme_session_review_step`` and are not included in this @@ -36,8 +36,7 @@ _VERDICT_SCHEMA = { "properties": { "reasoning": { "type": "string", - "description": "用中文写出详细的推理过程:先说明证据支持的答案,再逐步判断 golden_answer " - "是否正确,以及 answer_session_ids 是否恰好正确。", + "description": "用中文写出详细的推理过程:先说明证据支持的答案,再判断 " "golden_answer 是否正确。", }, "golden_answer_correct": { "type": "boolean", @@ -45,26 +44,14 @@ _VERDICT_SCHEMA = { }, "true_answer": { "type": "string", - "description": "仅当 golden_answer_correct 为 false 时填写:证据支持的正确答案(证据不足时填 " - "'unknown')。golden_answer_correct 为 true 时填空字符串。", - }, - "answer_session_ids_correct": { - "type": "boolean", - "description": "answer_session_ids 是否恰好是支持答案所需的会话(多、少、无关的 id 都算错误)。", - }, - "true_answer_session_ids": { - "type": "array", - "items": {"type": "string"}, - "description": "仅当 answer_session_ids_correct 为 false 时填写:真正支持答案的 session id 列表。" - "answer_session_ids_correct 为 true 时填空列表。", + "description": "仅当 golden_answer_correct 为 false 时填写:证据支持的正确答案" + "(证据不足时填 'unknown')。golden_answer_correct 为 true 时填空字符串。", }, }, "required": [ "reasoning", "golden_answer_correct", "true_answer", - "answer_session_ids_correct", - "true_answer_session_ids", ], "additionalProperties": False, } @@ -110,8 +97,6 @@ class GoldenCheckStep(BaseStep): question_type = str(query.get("question_type") or "").strip() question_date = str(query.get("question_date") or "").strip() golden_answer = str(golden.get("answer") or "").strip() - answer_session_ids = golden.get("answer_session_ids_filter_illegal") or [] - if not question: raise ValueError(f"{review_path} does not contain a question") @@ -120,7 +105,6 @@ class GoldenCheckStep(BaseStep): "question_type": question_type, "question_date": question_date, "golden_answer": golden_answer, - "answer_session_ids": answer_session_ids, "session_summaries": session_summaries, } user_prompt = self.prompt_format( @@ -129,7 +113,6 @@ class GoldenCheckStep(BaseStep): question_type=question_type, question_date=question_date, golden_answer=golden_answer, - answer_session_ids=", ".join(str(s) for s in answer_session_ids) or "(none)", num_session_summaries=len(session_summaries), payload_json=json.dumps(prompt_input, ensure_ascii=False, indent=2), ) @@ -173,6 +156,10 @@ class GoldenCheckStep(BaseStep): if not isinstance(verdict, dict): self.logger.warning(f"[{self.name}] no structured verdict; falling back to free text") verdict = {"reasoning": (result.get("result") or "").strip()} + # Retain the legacy fields for readers of existing check_golden.json + # artifacts. They are compatibility placeholders, not audit results. + verdict["answer_session_ids_correct"] = True + verdict["true_answer_session_ids"] = [] # Slim output: do NOT duplicate session_review.json (referenced by path); # keep only the compact session_summaries and the verdict. diff --git a/reme/steps/benchmark/lme/golden_check.yaml b/reme/steps/benchmark/lme/golden_check.yaml index d6b71d11..0bb590ba 100644 --- a/reme/steps/benchmark/lme/golden_check.yaml +++ b/reme/steps/benchmark/lme/golden_check.yaml @@ -1,10 +1,9 @@ system_prompt: | 你是 LongMemEval 基准测试的审核员。你要根据从用户聊天记录中提取的证据,判断某个问题的 - golden_answer 是否正确,以及 answer_session_ids 是否恰好正确。 + golden_answer 是否正确。只审核答案内容,不检查或评价 answer_session_ids。 - 你可以使用 python_execute 工具作为推理草稿本:统计相关会话、抽取答案、交叉核对 - answer_session_ids。把给定的数据以字面量形式直接嵌入 Python 代码,使计算可复现;把中间结果 - 以 JSON 打印出来,便于审计。 + 你可以使用 python_execute 工具作为推理草稿本:统计相关会话、抽取答案。把给定的数据以字面量 + 形式直接嵌入 Python 代码,使计算可复现;把中间结果以 JSON 打印出来,便于审计。 python 的 stdout 不是你的最终答案,只是草稿。计算充分、确信之后,停止调用 python,用中文 给出结论。最终的结构化结果会从整段对话中自动抽取,所以务必把推理和结论清楚表达。 @@ -12,20 +11,17 @@ system_prompt: | 结构化输出要求: - reasoning 必须是详细的中文推理过程。 - true_answer 仅在 golden_answer_correct 为 false 时填写,为空字符串否则。 - - true_answer_session_ids 仅在 answer_session_ids_correct 为 false 时填写,为空列表否则。 user_message: | Question: {question} Question type: {question_type} Question date: {question_date} Golden answer: {golden_answer} - Answer session ids: {answer_session_ids} Number of session extractions included: {num_session_summaries} - 证据(JSON)。下方 answer_session_ids 已是 session_review.json 中的 - answer_session_ids_filter_illegal;session_summaries 含上游审核过的会话,每条只有 session_id、 - session_date、extracted_info: + 证据(JSON)。session_summaries 含上游审核过的会话,每条只有 session_id、session_date、 + extracted_info: {payload_json} 用 python_execute 统计和推理,过程中打印中间 JSON。然后用中文给出最终结论:golden_answer - 是否正确(不正确时给出 true_answer),answer_session_ids 是否正确(不正确时给出真正的 id)。 + 是否正确(不正确时给出 true_answer)。不要评价 answer_session_ids。 diff --git a/reme/steps/benchmark/lme/session_review.yaml b/reme/steps/benchmark/lme/session_review.yaml index b891ace9..4605f407 100644 --- a/reme/steps/benchmark/lme/session_review.yaml +++ b/reme/steps/benchmark/lme/session_review.yaml @@ -10,8 +10,7 @@ system_prompt: | - Keep time expressions inline with the fact they modify, including dates, weekdays, relative times such as "last week" or "since January 15th", durations, and frequencies. - Do not invent facts. Only extract what is actually present in the session. - - Do not judge whether answer_session_ids are correct. That is handled by the downstream - golden check. + - Do not judge whether answer_session_ids are correct. They are outside the answer audit scope. - Output only the extracted information as plain text. Do not output JSON, markdown fences, or relevance labels. diff --git a/tests/unit/test_config_parser.py b/tests/unit/test_config_parser.py index 95838397..e056b433 100644 --- a/tests/unit/test_config_parser.py +++ b/tests/unit/test_config_parser.py @@ -49,17 +49,6 @@ def test_default_config_registers_daily_write_job(): assert job["parameters"]["required"] == ["name", "description", "session_id", "content"] -def test_default_config_registers_shell_job(): - """``shell`` exposes command execution through ``shell_step``.""" - cfg = _load_config("default.yaml") - - job = cfg["jobs"]["shell"] - assert job["backend"] == "base" - assert job["steps"] == [{"backend": "shell_step"}] - assert job["parameters"]["required"] == ["cmd"] - assert job["parameters"]["properties"]["shell_timeout"]["default"] == 86400 - - def test_default_config_keeps_frontmatter_chunk_metadata_opt_in(): """Markdown frontmatter-to-chunk metadata is disabled by default for compatibility.""" cfg = _load_config("default.yaml") diff --git a/tests/unit/test_lme_final_answer_review.py b/tests/unit/test_lme_final_answer_review.py new file mode 100644 index 00000000..65bb19e2 --- /dev/null +++ b/tests/unit/test_lme_final_answer_review.py @@ -0,0 +1,425 @@ +"""Focused tests for the disputed LongMemEval final-answer workflow.""" + +import asyncio +import json +import sys +import threading +import time +from pathlib import Path +from unittest.mock import AsyncMock + +import pytest + +from benchmark.longmemeval import run_final_answer_review as driver_module +from benchmark.longmemeval.run_final_answer_review import ( + REFERENCE_PATHS_ENV, + atomic_write_results, + merge_references, + select_question_ids, +) +from reme.components.agent_wrapper.base_agent_wrapper import BaseAgentWrapper +from reme.components.agent_wrapper.cc_agent_wrapper import CcAgentWrapper +from reme.components.application_context import ApplicationContext +from reme.config import resolve_app_config +from reme.steps.benchmark.lme import final_answer_review as review_module +from reme.steps.benchmark.lme.final_answer_review import FinalAnswerReviewStep + + +class _FakeAgentWrapper(BaseAgentWrapper): + """Return queued ordinary text replies and retain every prompt call.""" + + def __init__(self, replies: list[str]): + super().__init__() + self.replies = list(replies) + self.calls: list[tuple[str, dict]] = [] + + async def reply(self, inputs, **kwargs) -> dict: + """Return the next queued agent response.""" + self.calls.append((inputs, kwargs)) + return { + "session_id": f"attempt-{len(self.calls)}", + "result": self.replies.pop(0), + } + + +def _write_json(path: Path, value: object) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(value, ensure_ascii=False), encoding="utf-8") + + +def _write_jsonl(path: Path, rows: list[dict]) -> None: + path.write_text( + "".join(json.dumps(row, ensure_ascii=False) + "\n" for row in rows), + encoding="utf-8", + ) + + +def _session(session_id: str, date: str, marker: str) -> dict: + return { + "haystack_session_id": session_id, + "haystack_date": date, + "messages": [{"role": "user", "content": marker}], + "other_session_field": f"full-{marker}", + } + + +def test_final_answer_review_keeps_raw_sessions_out_of_prompt_and_retries_plain_json( + tmp_path, + monkeypatch, +): + """Raw session messages stay on disk, and invalid ordinary replies are retried.""" + query = { + "question_id": "question-1", + "question": "What happened?", + "question_type": "single-session-user", + "question_date": "2024/01/02 (Tue) 10:00", + "extra_query_field": "keep-me", + } + golden = { + "answer": "old answer", + "answer_session_ids": ["past", "future"], + "extra_answer_field": "keep-me-too", + } + _write_json(tmp_path / "query.json", query) + _write_json(tmp_path / "answer.json", golden) + _write_json( + tmp_path / "session" / "past.json", + _session("past", "2024/01/02 (Tue) 09:59", "past-evidence"), + ) + _write_json( + tmp_path / "session" / "equal.json", + _session("equal", "2024/01/02 (Tue) 10:00", "equal-evidence"), + ) + _write_json( + tmp_path / "session" / "future.json", + _session("future", "2024/01/02 (Tue) 10:01", "future-secret"), + ) + _write_jsonl( + tmp_path / "first.jsonl", + [ + { + "question_id": "question-1", + "answer": "reference one", + "reason": "first reason", + }, + ], + ) + _write_jsonl( + tmp_path / "second.jsonl", + [ + { + "question_id": "question-1", + "answer": "reference two", + "reason": "second reason", + }, + ], + ) + + wrapper = _FakeAgentWrapper( + [ + '{"reason":"missing fence","golden_answer_correct":false,"answer":"invalid",' + '"is_session_time_wrong":false}', + '```json\n{"reason":"deprecated timestamp verdict","golden_answer_correct":false,' + '"answer":"still invalid","is_session_time_wrong":true}\n```', + "补充分析可以放在代码块外。\n" + '```json\n{"reason":"由 past 和 equal 两个 session 支持 golden answer。",' + '"golden_answer_correct":true,"answer":"","is_session_time_wrong":false}\n```\n' + "审核完成。", + ], + ) + sleep = AsyncMock() + monkeypatch.setattr(review_module.asyncio, "sleep", sleep) + app_context = ApplicationContext( + workspace_dir=str(tmp_path), + resource_dir="session", + ) + step = FinalAnswerReviewStep( + app_context=app_context, + agent_wrapper=wrapper, + reference_paths=["first.jsonl", "second.jsonl"], + retry_initial_seconds=0.01, + retry_max_seconds=0.02, + ) + + response = asyncio.run(step()) + + assert response.success is True + assert json.loads(response.answer) == { + "reason": "由 past 和 equal 两个 session 支持 golden answer。", + "golden_answer_correct": True, + "answer": "", + "is_session_time_wrong": False, + } + assert response.metadata["attempts"] == 3 + assert response.metadata["num_sessions"] == 3 + assert response.metadata["num_future_sessions"] == 1 + assert response.metadata["future_sessions"] == [ + { + "session_id": "future", + "session_date": "2024/01/02 (Tue) 10:01", + "session_file": "future.json", + }, + ] + assert len(wrapper.calls) == 3 + prompt, reply_kwargs = wrapper.calls[0] + assert "past-evidence" not in prompt + assert "equal-evidence" not in prompt + assert "full-past-evidence" not in prompt + assert "future-secret" not in prompt + assert "extra_query_field" in prompt + assert "extra_answer_field" in prompt + assert "reference one" in prompt and "reference two" in prompt + assert '"session_time_check"' in prompt + assert '"sessions_after_question_date": [' in prompt + assert '"answer_session_ids_after_question_date"' not in prompt + assert '"future"' in prompt + assert "output_schema" not in reply_kwargs + assert [call.args for call in sleep.await_args_list] == [(0.01,), (0.02,)] + + +# pylint: disable=protected-access +def test_final_answer_review_reference_paths_env_overrides_config(tmp_path, monkeypatch): + """The batch driver can pass its selected reference files into the job process.""" + configured = tmp_path / "configured.jsonl" + selected = tmp_path / "selected.jsonl" + _write_jsonl( + configured, + [{"question_id": "question-1", "answer": "configured", "reason": "configured reason"}], + ) + _write_jsonl( + selected, + [{"question_id": "question-1", "answer": "selected", "reason": "selected reason"}], + ) + monkeypatch.setenv(REFERENCE_PATHS_ENV, json.dumps([str(selected)])) + step = FinalAnswerReviewStep( + app_context=ApplicationContext(workspace_dir=str(tmp_path)), + reference_paths=[str(configured)], + ) + + references = step._load_references("question-1") + + assert len(references) == 1 + assert references[0]["answer"] == "selected" + assert references[0]["source"] == selected.name + + +def test_final_answer_review_allows_question_without_reference_answer(tmp_path): + """Samples outside the disputed lists are reviewed from answer.json alone.""" + references_path = tmp_path / "references.jsonl" + _write_jsonl( + references_path, + [{"question_id": "another-question", "answer": "other", "reason": "other reason"}], + ) + step = FinalAnswerReviewStep( + app_context=ApplicationContext(workspace_dir=str(tmp_path)), + reference_paths=[str(references_path)], + ) + + assert not step._load_references("question-without-reference") + + +# pylint: enable=protected-access + + +def test_final_answer_review_agent_cwd_is_sample_session_directory(tmp_path): + """The configured relative cwd resolves inside each selected LME workspace.""" + config = resolve_app_config(config="jinli_lme", log_config=False) + agent_config = config["components"]["agent_wrapper"]["lme_final_answer_review"] + assert agent_config["cwd"] == "session" + + wrapper = CcAgentWrapper( + app_context=ApplicationContext(workspace_dir=str(tmp_path)), + cwd=agent_config["cwd"], + ) + assert wrapper.cwd == tmp_path / "session" + + +# pylint: disable=protected-access +def test_final_answer_review_requires_empty_answer_when_golden_is_correct(): + """Correct golden answers are collected without duplicating their answer text.""" + parsed = FinalAnswerReviewStep._parse_reply( + '```json\n{"reason":"golden is supported","golden_answer_correct":true,"answer":"",' + '"is_session_time_wrong":false}\n```', + ) + assert parsed == { + "reason": "golden is supported", + "golden_answer_correct": True, + "answer": "", + "is_session_time_wrong": False, + } + + with pytest.raises(ValueError, match="answer.*must be empty"): + FinalAnswerReviewStep._parse_reply( + '```json\n{"reason":"bad duplicate","golden_answer_correct":true,"answer":"duplicate",' + '"is_session_time_wrong":false}\n```', + ) + + with pytest.raises(ValueError, match="deprecated and must be false"): + FinalAnswerReviewStep._parse_reply( + '```json\n{"reason":"legacy session id verdict",' + '"golden_answer_correct":false,"answer":"corrected","is_session_time_wrong":true}\n```', + ) + + with pytest.raises(ValueError, match="must not evaluate answer_session_ids"): + FinalAnswerReviewStep._parse_reply( + '```json\n{"reason":"answer_session_ids contains a future session",' + '"golden_answer_correct":false,"answer":"corrected","is_session_time_wrong":false}\n```', + ) + + +# pylint: enable=protected-access + + +def test_final_answer_review_rejects_unparseable_session_time_before_agent(tmp_path): + """An unknown session time is never silently admitted across the time boundary.""" + _write_json( + tmp_path / "query.json", + { + "question_id": "question-1", + "question": "Q", + "question_date": "2024/01/02 (Tue) 10:00", + }, + ) + _write_json(tmp_path / "answer.json", {"answer": "A"}) + _write_json( + tmp_path / "session" / "bad.json", + _session("bad", "unknown", "must-not-reach-agent"), + ) + _write_jsonl( + tmp_path / "refs.jsonl", + [{"question_id": "question-1", "answer": "reference", "reason": "reason"}], + ) + valid_reply = "".join( + [ + '```json\n{"reason":"y","golden_answer_correct":false,', + '"answer":"x","is_session_time_wrong":false}\n```', + ], + ) + wrapper = _FakeAgentWrapper([valid_reply]) + step = FinalAnswerReviewStep( + app_context=ApplicationContext( + workspace_dir=str(tmp_path), + resource_dir="session", + ), + agent_wrapper=wrapper, + reference_paths=["refs.jsonl"], + ) + + with pytest.raises(ValueError, match="Invalid LongMemEval datetime"): + asyncio.run(step()) + assert not wrapper.calls + + +def test_driver_merges_references_and_atomically_rewrites_in_input_order(tmp_path): + """The batch checkpoint contains one stable row per completed question.""" + first = tmp_path / "first.jsonl" + second = tmp_path / "second.jsonl" + _write_jsonl( + first, + [ + {"question_id": "q2", "answer": "a2", "reason": "r2"}, + {"question_id": "q1", "answer": "a1", "reason": "r1"}, + ], + ) + _write_jsonl(second, [{"question_id": "q1", "answer": "a1b", "reason": "r1b"}]) + + merged = merge_references([first, second]) + + assert list(merged) == ["q2", "q1"] + assert len(merged["q2"]) == 1 + assert len(merged["q1"]) == 2 + output = tmp_path / "result.jsonl" + atomic_write_results( + output, + list(merged), + { + "q1": { + "reason": "reason-1", + "golden_answer_correct": False, + "answer": "final-1", + "is_session_time_wrong": False, + }, + "q2": { + "reason": "reason-2", + "golden_answer_correct": False, + "answer": "final-2", + "is_session_time_wrong": True, + }, + }, + ) + rows = _read_output(output) + assert [row["question_id"] for row in rows] == ["q2", "q1"] + assert driver_module.load_existing(output)["q2"]["is_session_time_wrong"] is False + + +def test_driver_selects_all_or_explicit_question_ids(tmp_path): + """Explicit IDs may select samples that have no reference-answer row.""" + mapping = { + "q1": tmp_path / "0", + "q2": tmp_path / "1", + "q3": tmp_path / "2", + } + + assert select_question_ids(mapping, None) == ["q1", "q2", "q3"] + assert select_question_ids(mapping, ["q3", "q1"]) == ["q3", "q1"] + assert select_question_ids(mapping, None, {"q1", "q3"}) == ["q2"] + assert select_question_ids(mapping, ["q3", "q2"], {"q3"}) == ["q2"] + with pytest.raises(ValueError, match="No dataset workspace"): + select_question_ids(mapping, ["unknown"]) + with pytest.raises(ValueError, match="Duplicate"): + select_question_ids(mapping, ["q1", "q1"]) + + +def test_driver_limits_concurrency_and_spaces_submissions(tmp_path, monkeypatch): + """Concurrent jobs never exceed the cap and are not submitted in a burst.""" + mapping = {f"q{index}": tmp_path / str(index) for index in range(4)} + starts: list[float] = [] + active = 0 + max_active = 0 + lock = threading.Lock() + + def fake_run_one(question_id, workspace, log_dir, reference_paths): + del question_id, workspace, log_dir, reference_paths + nonlocal active, max_active + with lock: + starts.append(time.monotonic()) + active += 1 + max_active = max(max_active, active) + time.sleep(0.055) + with lock: + active -= 1 + return { + "reason": "reviewed", + "golden_answer_correct": True, + "answer": "", + "is_session_time_wrong": False, + } + + monkeypatch.setattr(driver_module, "workspace_map", lambda: mapping) + monkeypatch.setattr(driver_module, "merge_references", lambda paths: {}) + monkeypatch.setattr(driver_module, "load_existing", lambda path: {}) + monkeypatch.setattr(driver_module, "atomic_write_results", lambda *args: None) + monkeypatch.setattr(driver_module, "run_one", fake_run_one) + monkeypatch.setattr(driver_module, "MIN_SUBMIT_INTERVAL_SECONDS", 0.0) + monkeypatch.setattr( + sys, + "argv", + [ + "run_final_answer_review.py", + "--concurrency", + "3", + "--submit-interval-seconds", + "0.02", + "--output", + str(tmp_path / "output.jsonl"), + ], + ) + + assert driver_module.main() == 0 + assert max_active == 3 + assert len(starts) == 4 + assert all(later - earlier >= 0.015 for earlier, later in zip(starts, starts[1:])) + + +def _read_output(path: Path) -> list[dict]: + return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines()]