feat: add eval corpus data generator

7-phase pipeline for generating synthetic multi-file corpora:
1. Scenario Brief (SCENARIO.md) - world-building
2. Fact Registry (facts.json) - consistency source of truth
3. File Manifest (manifest.json) - per-file briefs
4. Clustering - topological sort + fact registry sharding
5. Parallel File Generation - concurrent workers
6. Validation - cross-reference & consistency audit
7. Question Generation - 10 eval questions per corpus

Supports all 13 data points (dp_001 through dp_013, 5 to 10,000 files).
Uses Gemini 2.5 Pro by default. Includes resume support, validation,
and 219 unit tests.
This commit is contained in:
Dhravya 2026-04-28 23:24:42 +00:00
parent 0d2ca1bb41
commit cba994be3f
19 changed files with 7847 additions and 0 deletions

107
data-generator/README.md Normal file
View file

@ -0,0 +1,107 @@
# Eval Corpus Data Generator
Generates synthetic multi-file corpora for the SMFS memory eval benchmark. Each corpus simulates a real organization's shared memory — files written by many authors, in many formats, over a specific time period.
## Architecture
7-phase pipeline:
1. **Scenario Brief** — One LLM call creates the "bible" for the corpus (cast, timeline, locked facts, per-file briefs)
2. **Fact Registry** — Extracts every concrete fact into structured JSON (the single source of truth for consistency)
3. **File Manifest** — Describes every file to generate (path, format, author, locked facts, cross-references)
4. **Clustering** — Groups files into clusters of 3-8, topologically sorted so dependencies generate first
5. **File Generation** — Parallel workers generate files within clusters, passing cross-reference context
6. **Validation** — Audits token counts, locked facts, name consistency, cross-references
7. **Question Generation** — Creates 10 eval questions per corpus
## Setup
```bash
pip install -r requirements.txt
```
Requires a Gemini API key (or other LLM provider key) in the environment:
```bash
export GEMINI_API_KEY=your-key-here
# or
export OPENAI_API_KEY=your-key-here
export ANTHROPIC_API_KEY=your-key-here
```
## Usage
```bash
# Generate a single data point
python generate.py dp_001
# Generate a range
python generate.py dp_001 dp_005
# Resume a failed generation
python generate.py dp_003 --resume
# Generate only questions for an existing corpus
python generate.py dp_001 --questions-only
# Validate an existing corpus
python generate.py dp_002 --validate-only
# Use a specific model
python generate.py dp_001 --model openai/gpt-4o
# Set concurrency for large corpora
python generate.py dp_010 --max-concurrent 20
# Custom output directory
python generate.py dp_001 --output-dir /path/to/output
```
## Data Points
| dp | files | scenario |
|----|-------|----------|
| dp_001 | 5 | Two-person consulting kickoff |
| dp_002 | 10 | Couple's anniversary weekend trip |
| dp_003 | 20 | Single ER patient case across visits |
| dp_004 | 30 | Small-claims legal matter |
| dp_005 | 50 | Two-roommate co-living journal |
| dp_006 | 100 | Indie open-source project, 6 months |
| dp_007 | 200 | Grad-student lab, first semester |
| dp_008 | 300 | Pre-seed startup, first 6 months |
| dp_009 | 500 | Small therapy practice, 6 months |
| dp_010 | 1,000 | Growth-stage startup, 6 months |
| dp_011 | 2,000 | Newsroom investigation, 18 months |
| dp_012 | 5,000 | Embassy at one posting, 3-year archive |
| dp_013 | 10,000 | Tech-company CEO, full annual archive |
## Output Structure
```
output/dp_NNN/
├── SCENARIO.md # Deep brief (world-building bible)
├── facts.json # Structured fact registry
├── manifest.json # File manifest with per-file briefs
├── data/ # The actual corpus
│ ├── [domain folders]/
│ └── memory/
│ ├── profiles/
│ └── ...
├── question.json # 10 eval questions
├── generation_log.json # Audit trail (model, tokens, retries)
└── validation_report.json # Consistency audit results
```
## Testing
```bash
python -m pytest test_planner.py test_clusterer.py test_worker.py test_validator.py test_questions.py -v
```
## Design Decisions
- **Gemini 2.5 Pro** as default model (free tier, large context window)
- **Fact registry sharding** for large corpora: global facts (people, orgs) go to every worker; scoped facts (dates, financials) only go to workers that need them
- **Topological cluster ordering**: files that cross-reference each other are co-generated; dependency clusters generate first
- **30% overshoot tolerance** on token counts: slightly long is better than too short
- **Resume support**: every phase checks for existing output and skips if found

546
data-generator/clusterer.py Normal file
View file

@ -0,0 +1,546 @@
"""Phase 4: Cluster assignment with topological sort and fact registry sharding.
Takes a file manifest and a fact registry, groups files into generation clusters,
topologically sorts them so dependencies generate first, and shards the fact
registry so each worker receives only the facts it needs.
"""
from __future__ import annotations
import logging
import warnings
from collections import defaultdict, deque
from dataclasses import dataclass, field
from typing import Any
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Public data types
# ---------------------------------------------------------------------------
MIN_CLUSTER_SIZE = 1
DEFAULT_MAX_CLUSTER_SIZE = 8
# Categories in the fact registry that are scoped to specific files.
_SCOPED_CATEGORIES = ("dates", "financial", "references", "locations", "domain_facts")
# Categories that are global — every worker needs these.
_GLOBAL_CATEGORIES = ("people", "organizations")
@dataclass
class Cluster:
"""A group of related files to be generated together."""
cluster_id: str
file_entries: list[dict] = field(default_factory=list)
fact_shard: dict = field(default_factory=dict)
depends_on: list[str] = field(default_factory=list)
level: int = 0
# ---------------------------------------------------------------------------
# Fact registry sharding
# ---------------------------------------------------------------------------
def shard_fact_registry(fact_registry: dict, cluster_file_ids: list[str]) -> dict:
"""Return a subset of the fact registry relevant to the given file IDs.
Always includes: all people, all organizations (these are global).
Filters: dates, financial, references, locations, domain_facts only those
whose 'files' array intersects with cluster_file_ids.
Cross_references: only those where source or target is in cluster_file_ids.
"""
file_id_set = set(cluster_file_ids)
shard: dict[str, Any] = {}
# Copy top-level scalar fields (e.g. scenario_id)
for key, value in fact_registry.items():
if not isinstance(value, list):
shard[key] = value
# Global categories — always included in full
for category in _GLOBAL_CATEGORIES:
if category in fact_registry:
shard[category] = list(fact_registry[category])
# Scoped categories — filter to entries whose files intersect
for category in _SCOPED_CATEGORIES:
if category not in fact_registry:
continue
filtered = [
entry
for entry in fact_registry[category]
if file_id_set.intersection(entry.get("files", []))
]
shard[category] = filtered
# Cross-references — keep only those touching our files
if "cross_references" in fact_registry:
shard["cross_references"] = [
xref
for xref in fact_registry["cross_references"]
if xref.get("source_file") in file_id_set
or xref.get("target_file") in file_id_set
]
return shard
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _get_file_id(entry: dict) -> str:
"""Extract the file_id from a manifest entry."""
return entry.get("file_id", "")
def _get_cluster_hint(entry: dict) -> str:
"""Extract the cluster_hint from a manifest entry, defaulting to 'misc'."""
return entry.get("cluster_hint", "misc") or "misc"
def _get_cross_references(entry: dict) -> list[str]:
"""Extract cross_references from a manifest entry."""
refs = entry.get("cross_references", [])
if isinstance(refs, list):
return refs
return []
def _build_file_id_to_entry(manifest: list[dict]) -> dict[str, dict]:
"""Build a lookup from file_id to manifest entry."""
return {_get_file_id(e): e for e in manifest if _get_file_id(e)}
def _group_by_cluster_hint(manifest: list[dict]) -> dict[str, list[dict]]:
"""Group manifest entries by their cluster_hint field."""
groups: dict[str, list[dict]] = defaultdict(list)
for entry in manifest:
hint = _get_cluster_hint(entry)
groups[hint].append(entry)
return dict(groups)
def _build_cross_ref_graph(
manifest: list[dict], valid_file_ids: set[str]
) -> dict[str, set[str]]:
"""Build an adjacency list of cross-references between files.
Returns a mapping from file_id -> set of file_ids it references.
Warns and ignores references to files not in the manifest.
"""
graph: dict[str, set[str]] = defaultdict(set)
for entry in manifest:
fid = _get_file_id(entry)
for ref in _get_cross_references(entry):
if ref not in valid_file_ids:
warnings.warn(
f"File '{fid}' cross-references '{ref}' which is not in the manifest; ignoring.",
stacklevel=2,
)
continue
if ref != fid:
graph[fid].add(ref)
graph[ref].add(fid) # bidirectional for clustering purposes
return dict(graph)
def _find_connected_components(
file_ids: list[str], adjacency: dict[str, set[str]]
) -> list[list[str]]:
"""Find connected components within a set of file_ids using the adjacency graph."""
id_set = set(file_ids)
visited: set[str] = set()
components: list[list[str]] = []
for fid in file_ids:
if fid in visited:
continue
component: list[str] = []
queue = deque([fid])
while queue:
current = queue.popleft()
if current in visited or current not in id_set:
continue
visited.add(current)
component.append(current)
for neighbor in adjacency.get(current, set()):
if neighbor in id_set and neighbor not in visited:
queue.append(neighbor)
if component:
components.append(component)
return components
def _split_group(
entries: list[dict],
max_size: int,
adjacency: dict[str, set[str]],
) -> list[list[dict]]:
"""Split a group that exceeds max_size into smaller chunks.
Keeps cross-referencing files together where possible.
"""
file_ids = [_get_file_id(e) for e in entries]
id_to_entry = {_get_file_id(e): e for e in entries}
# Find connected components within this group
components = _find_connected_components(file_ids, adjacency)
# Build sub-groups, packing components into chunks up to max_size
sub_groups: list[list[dict]] = []
current: list[dict] = []
for component in components:
component_entries = [id_to_entry[fid] for fid in component]
if len(component_entries) > max_size:
# Component itself is too big — forcibly split it
if current:
sub_groups.append(current)
current = []
for i in range(0, len(component_entries), max_size):
sub_groups.append(component_entries[i : i + max_size])
elif len(current) + len(component_entries) > max_size:
# Adding this component would exceed limit — start a new chunk
if current:
sub_groups.append(current)
current = list(component_entries)
else:
current.extend(component_entries)
if current:
sub_groups.append(current)
return sub_groups
def _try_merge_singletons(
groups: dict[str, list[dict]],
max_size: int,
adjacency: dict[str, set[str]],
) -> dict[str, list[dict]]:
"""Merge singleton groups into a group they cross-reference, if room allows."""
singleton_keys = [k for k, v in groups.items() if len(v) == 1]
file_to_group: dict[str, str] = {}
for gkey, entries in groups.items():
for entry in entries:
file_to_group[_get_file_id(entry)] = gkey
merged_into: dict[str, str] = {} # singleton_key -> target_key
for skey in singleton_keys:
entry = groups[skey][0]
fid = _get_file_id(entry)
refs = adjacency.get(fid, set())
for ref in refs:
target_group = file_to_group.get(ref)
if (
target_group
and target_group != skey
and target_group not in merged_into.values() # don't chain-merge
and len(groups[target_group]) < max_size
):
groups[target_group].append(entry)
file_to_group[fid] = target_group
merged_into[skey] = target_group
break
for skey in merged_into:
del groups[skey]
return groups
def _build_cluster_dependency_graph(
clusters: dict[str, list[dict]],
file_to_cluster: dict[str, str],
manifest: list[dict],
valid_file_ids: set[str],
) -> dict[str, set[str]]:
"""Build a DAG of cluster dependencies from cross-references.
If a file in cluster A references a file in cluster B (and A != B),
then A depends on B (B must generate before A).
Returns: mapping from cluster_id -> set of cluster_ids it depends on.
"""
deps: dict[str, set[str]] = defaultdict(set)
for entry in manifest:
fid = _get_file_id(entry)
source_cluster = file_to_cluster.get(fid)
if not source_cluster:
continue
for ref in _get_cross_references(entry):
if ref not in valid_file_ids:
continue
target_cluster = file_to_cluster.get(ref)
if target_cluster and target_cluster != source_cluster:
deps[source_cluster].add(target_cluster)
return dict(deps)
def _detect_and_merge_cycles(
cluster_groups: dict[str, list[dict]],
deps: dict[str, set[str]],
) -> tuple[dict[str, list[dict]], dict[str, set[str]]]:
"""Detect cycles in the dependency graph and merge cyclic clusters.
Uses Tarjan-like SCC detection via iterative DFS.
Returns updated cluster_groups and deps with cycles removed.
"""
all_ids = set(cluster_groups.keys())
# Iterative Tarjan's SCC algorithm
index_counter = [0]
stack: list[str] = []
on_stack: set[str] = set()
indices: dict[str, int] = {}
lowlinks: dict[str, int] = {}
sccs: list[list[str]] = []
def strongconnect(v: str) -> None:
# Iterative version using explicit call stack
call_stack: list[tuple[str, list[str], int]] = []
indices[v] = lowlinks[v] = index_counter[0]
index_counter[0] += 1
stack.append(v)
on_stack.add(v)
neighbors = sorted(deps.get(v, set()) & all_ids)
call_stack.append((v, neighbors, 0))
while call_stack:
node, nbrs, idx = call_stack[-1]
if idx < len(nbrs):
call_stack[-1] = (node, nbrs, idx + 1)
w = nbrs[idx]
if w not in indices:
indices[w] = lowlinks[w] = index_counter[0]
index_counter[0] += 1
stack.append(w)
on_stack.add(w)
w_neighbors = sorted(deps.get(w, set()) & all_ids)
call_stack.append((w, w_neighbors, 0))
elif w in on_stack:
lowlinks[node] = min(lowlinks[node], indices[w])
else:
# All neighbors processed
if lowlinks[node] == indices[node]:
scc: list[str] = []
while True:
w = stack.pop()
on_stack.discard(w)
scc.append(w)
if w == node:
break
sccs.append(scc)
call_stack.pop()
if call_stack:
parent = call_stack[-1][0]
lowlinks[parent] = min(lowlinks[parent], lowlinks[node])
for cid in sorted(all_ids):
if cid not in indices:
strongconnect(cid)
# Merge SCCs with more than one node
merge_map: dict[str, str] = {} # old_id -> merged_id
for scc in sccs:
if len(scc) <= 1:
continue
# Merge all into the first (alphabetically sorted)
scc_sorted = sorted(scc)
primary = scc_sorted[0]
for cid in scc_sorted[1:]:
merge_map[cid] = primary
cluster_groups[primary].extend(cluster_groups.pop(cid))
logger.info(
f"Merged cyclic clusters {scc_sorted[1:]} into '{primary}'"
)
if not merge_map:
return cluster_groups, deps
# Rebuild dependency graph with merged IDs
new_deps: dict[str, set[str]] = defaultdict(set)
for cid, dep_set in deps.items():
resolved_cid = merge_map.get(cid, cid)
if resolved_cid not in cluster_groups:
continue
for d in dep_set:
resolved_d = merge_map.get(d, d)
if resolved_d != resolved_cid and resolved_d in cluster_groups:
new_deps[resolved_cid].add(resolved_d)
return cluster_groups, dict(new_deps)
def _topological_sort_with_levels(
cluster_ids: list[str],
deps: dict[str, set[str]],
) -> list[tuple[str, int]]:
"""Kahn's algorithm producing (cluster_id, level) pairs.
Level 0 = no dependencies. Level N = max dependency level + 1.
Returns pairs sorted by level, then cluster_id.
"""
all_ids = set(cluster_ids)
# Build in-degree and adjacency
in_degree: dict[str, int] = {cid: 0 for cid in all_ids}
# forward edges: dep -> [dependents]
forward: dict[str, list[str]] = defaultdict(list)
for cid in all_ids:
for dep in deps.get(cid, set()):
if dep in all_ids:
in_degree[cid] += 1
forward[dep].append(cid)
# Initialize queue with all nodes that have in-degree 0
queue: deque[str] = deque()
levels: dict[str, int] = {}
for cid in sorted(all_ids):
if in_degree[cid] == 0:
queue.append(cid)
levels[cid] = 0
result: list[tuple[str, int]] = []
while queue:
current = queue.popleft()
result.append((current, levels[current]))
for dependent in forward.get(current, []):
in_degree[dependent] -= 1
levels[dependent] = max(
levels.get(dependent, 0), levels[current] + 1
)
if in_degree[dependent] == 0:
queue.append(dependent)
# Safety check: if we didn't visit all nodes, there's an unexpected cycle
if len(result) < len(all_ids):
missing = all_ids - {cid for cid, _ in result}
logger.warning(
f"Topological sort did not visit all clusters. "
f"Remaining (possible cycle): {missing}. Assigning max level."
)
max_level = max((lvl for _, lvl in result), default=0) + 1
for cid in sorted(missing):
result.append((cid, max_level))
result.sort(key=lambda pair: (pair[1], pair[0]))
return result
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def assign_clusters(
manifest: list[dict],
fact_registry: dict,
max_cluster_size: int = DEFAULT_MAX_CLUSTER_SIZE,
) -> list[Cluster]:
"""Group files into clusters, topologically sort, and shard fact registry.
Returns clusters ordered by level (level 0 first, then level 1, etc.).
Within a level, clusters can run in parallel.
Algorithm:
1. Group files by cluster_hint from manifest
2. Split oversize groups (keeping cross-referencing files together)
3. Merge singletons into groups they cross-reference if room allows
4. Build cross-cluster dependency graph from cross_references
5. Detect and merge cycles
6. Topological sort using Kahn's algorithm
7. Assign levels (distance from root in the DAG)
8. Shard fact registry for each cluster
"""
if not manifest:
return []
valid_file_ids = {_get_file_id(e) for e in manifest if _get_file_id(e)}
file_id_to_entry = _build_file_id_to_entry(manifest)
# Step 1: Group by cluster_hint
groups = _group_by_cluster_hint(manifest)
# Build cross-reference adjacency for splitting/merging decisions
adjacency = _build_cross_ref_graph(manifest, valid_file_ids)
# Step 2: Split oversize groups
split_groups: dict[str, list[dict]] = {}
counter = 0
for hint, entries in groups.items():
if len(entries) <= max_cluster_size:
split_groups[hint] = entries
else:
sub_groups = _split_group(entries, max_cluster_size, adjacency)
for i, sub in enumerate(sub_groups):
key = f"{hint}_{i}" if len(sub_groups) > 1 else hint
split_groups[key] = sub
counter += 1
# Step 3: Merge singletons
split_groups = _try_merge_singletons(split_groups, max_cluster_size, adjacency)
# Build file_to_cluster mapping
file_to_cluster: dict[str, str] = {}
for cid, entries in split_groups.items():
for entry in entries:
file_to_cluster[_get_file_id(entry)] = cid
# Step 4: Build cross-cluster dependency graph
deps = _build_cluster_dependency_graph(
split_groups, file_to_cluster, manifest, valid_file_ids
)
# Step 5: Detect and merge cycles
split_groups, deps = _detect_and_merge_cycles(split_groups, deps)
# Rebuild file_to_cluster after potential merges
file_to_cluster = {}
for cid, entries in split_groups.items():
for entry in entries:
file_to_cluster[_get_file_id(entry)] = cid
# Rebuild deps after merge (edges may have changed)
deps = _build_cluster_dependency_graph(
split_groups, file_to_cluster, manifest, valid_file_ids
)
# Steps 6-7: Topological sort with levels
sorted_pairs = _topological_sort_with_levels(
list(split_groups.keys()), deps
)
# Step 8: Build Cluster objects with sharded fact registries
clusters: list[Cluster] = []
for cluster_id, level in sorted_pairs:
entries = split_groups[cluster_id]
cluster_file_ids = [_get_file_id(e) for e in entries]
fact_shard = shard_fact_registry(fact_registry, cluster_file_ids)
depends_on = sorted(deps.get(cluster_id, set()))
clusters.append(
Cluster(
cluster_id=cluster_id,
file_entries=entries,
fact_shard=fact_shard,
depends_on=depends_on,
level=level,
)
)
return clusters

668
data-generator/generate.py Normal file
View file

@ -0,0 +1,668 @@
#!/usr/bin/env python3
"""CLI entry point for the eval corpus data generator.
Usage:
# Generate a single data point
python generate.py dp_001
# Generate a range of data points
python generate.py dp_001 dp_005
# Resume a failed generation
python generate.py dp_003 --resume
# Generate only questions for an existing corpus
python generate.py dp_001 --questions-only
# Validate an existing corpus
python generate.py dp_002 --validate-only
# Use a specific model
python generate.py dp_001 --model gemini/gemini-2.5-pro
# Set concurrency
python generate.py dp_006 --max-concurrent 20
"""
from __future__ import annotations
import argparse
import asyncio
import json
import logging
import sys
import time
from pathlib import Path
from clusterer import assign_clusters
from planner import run_planning
from questions import generate_questions
from utils import DEFAULT_MODEL, GenerationLog, read_json, read_text, write_json
from validator import validate_corpus
from worker import generate_all
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Scenario definitions
# ---------------------------------------------------------------------------
# Each dp maps to a scenario block (the text from the eval design doc) and a
# file count. The scenario_block is the full description that gets passed to
# the planner for world-building.
SCENARIOS: dict[str, dict] = {
"dp_001": {
"file_count": 5,
"scenario_block": """\
## dp_001 — Two-person consulting engagement, day one
- **Files:** 5
- **Realized as:** Orbital Data (boutique data engineering consultancy) signs Coppertide \
(DTC cookware brand). Day one of a 12-week analytics-modernization engagement.
- **Reference date:** 2026-04-22 (kickoff).
- **Querier:** Priya Iyer (co-founder of Orbital, lead consultant; Bangalore tz; \
ex-Stripe, ex-Square; vegetarian, peanut allergy).
- **Cast:** Priya Iyer, Marcus Lehrer (Orbital co-founder, Berlin tz). Coppertide: \
Devansh Mehta (CTO), Aria Tan (Head of Analytics), Quentin Reyes (data eng), \
Lina Costa (VP Marketing).
- **Directory shape:** `client/coppertide/`, `internal/orbital/`, `memory/profiles/`, \
`memory/companies/`.
- **File mix:** signed SoW, detailed kickoff-call transcript, internal engagement plan, \
Priya's persona profile, Coppertide company overview.
- **Eval stressors:** floor case (every surface should pass at 5 files), single-hop, \
multi-hop chaining SoW + kickoff + engagement plan, profile.md cheap-cat for \
vegetarian/allergy/Bangalore facts.
- **Notable cross-references:** Stitch $2,034/mo, Snowflake $800/mo, Looker $1,400/mo, \
Klaviyo 14M rows/day, Spring Pans campaign, Thursday 1 PM ET review cadence, addendum \
due 2026-04-29, 2024 Fivetran/Shopify Plus duplicate-orders memory, SOC 2 \
PII-stays-in-US requirement.
- **Deep brief:** files implicitly carry the brief; consider creating \
`dp_001/SCENARIO.md` retrospectively for parity with dp_002.
""",
},
"dp_002": {
"file_count": 10,
"scenario_block": """\
## dp_002 — Couple's anniversary weekend trip
- **Files:** 10
- **Realized as:** Ana Sokol + Jordan Lee, long weekend in Portsmouth, NH \
(FriSun 2026-03-27 to 2026-03-29). 5-year dating anniversary.
- **Reference date:** 2026-03-25 (two days before departure).
- **Querier:** Ana Sokol (UX designer at Murex Health, Brooklyn).
- **Cast:** Ana Sokol, Jordan Lee (school librarian, Greene Hill Charter), \
Mira Bhattacharya (college friend in Portsmouth), Tomas Hjelm (college friend \
in Kittery, ME, briefly dated in 2014), Carolyn + Paul Foley (Martin Hill Inn \
innkeepers), Bea Acharya (downstairs neighbor caring for the cat Cipher), \
Yusra Marin (Jordan's librarian colleague), Priya Kuznetsov (Ana's manager \
first-name collision with dp_001's Priya Iyer; flagged).
- **Directory shape:** `trip/{itinerary,bookings,email,messages,notes}/`, \
`memory/{profiles,places}/`.
- **File mix:** shared itinerary doc, hotel + Amtrak confirmation .eml files with \
threaded replies, restaurants tracking, Mira's recommendation email + Saturday-lunch \
pushback, Tomas's nervous coffee-meetup email, two-week iMessage thread, Jordan's \
separate must-do list, Ana's persona profile, Portsmouth destination overview.
- **Eval stressors:** format-spanning across .eml and .md, multi-hop (Mira pushback \
chains email + restaurants doc), profile.md (Ana's pescatarian-this-trip, reading \
habits, Cipher), edit-then-recall (RiverRun-Bookstore-as-gift secret).
- **Locked facts:** booking refs `AMTKB-9F2RT-3K`, `MHINN-2026-0327-AS`, hotel rate \
$315/night x 2 + $94.50 NH tax + $30 pet deposit waived = $724.50, Amtrak fare \
$478.40 to Visa-4187, Acela 2151 / Downeaster 685 / Downeaster 690 / Acela 2168, \
all named restaurants.
- **Deep brief:** `.scratch/eval/test/dp_002/SCENARIO.md`.
""",
},
"dp_003": {
"file_count": 20,
"scenario_block": """\
## dp_003 — Single ER patient case across visits
- **Files:** 20
- **Setting:** A single patient case, ER admission through a 3-week follow-up. \
Mid-size urban hospital. Multiple providers touching the chart.
- **Querier:** The attending physician (or, alternatively, the patient themselves \
looking back at their own care). Default: attending.
- **Cast:** ER attending, ER nurse, hospitalist (admit), specialists pulled in \
(e.g., cardiology, GI), discharge planner, outpatient PCP, the patient, possibly \
a family member on the contact log.
- **Directory shape:** `clinical/{admission,progress,specialty}/`, \
`tests/{labs,imaging,reports}/`, `correspondence/{patient,family,provider}/`, \
`memory/profiles/`, `memory/conditions/`.
- **File mix:** ER admission note, triage assessment, lab orders + results \
(sometimes structured), imaging report, specialty consults (2-3), discharge summary, \
prescription notes, follow-up plan, patient symptom journal, billing note, provider \
1:1 patient-handoff note, optional anonymized peer consult, profile of patient, \
condition reference note.
- **Eval stressors:** longitudinal across a 3-week window for a single entity, \
format-spanning (lab CSV-style + imaging PDF transcription + prose progress notes), \
multi-hop (test result -> specialty consult -> med change), edit-then-recall \
(provider notes a follow-up).
- **Realism notes:** anonymized; PHI-shaped formatting (DOB, MRN); HL7-ish lab \
snippets; honest medical jargon density.
""",
},
"dp_004": {
"file_count": 30,
"scenario_block": """\
## dp_004 — Small-claims legal matter, intake to first hearing
- **Files:** 30
- **Setting:** A solo or two-person law practice handling a small-claims matter \
from intake through the first hearing (~6-week span).
- **Querier:** The lead attorney.
- **Cast:** Lead attorney, paralegal, opposing counsel (one or two), client, court \
clerk, possibly a witness or two.
- **Directory shape:** `client_intake/`, `pleadings/`, \
`correspondence/{client,opposing,court}/`, `research/`, `notes/`, `memory/profiles/`.
- **File mix:** client intake form, retainer agreement, demand letter, court filings \
(complaint, answer, motions), discovery requests + responses, attorney research memos, \
client correspondence (email + texts), opposing counsel correspondence, court \
communications, hearing prep notes, attorney's running case file, profile.
- **Eval stressors:** citation chains within a case (which filing references which \
exhibit), formal-correspondence retrieval, multi-hop (client said X in intake -> \
demand letter cites X -> opposing's answer responds to X), format-spanning across \
legal-PDF-shape and prose memos.
- **Realism notes:** docket numbers, plausible jurisdiction (Delaware or NY \
small-claims), realistic hearing date and motion practice; IP boilerplate \
inappropriate here small-claims is brief.
""",
},
"dp_005": {
"file_count": 50,
"scenario_block": """\
## dp_005 — Two-roommate co-living journal
- **Files:** 50
- **Setting:** Two roommates (not a couple) sharing an apartment for ~2 months. \
A shared journal, bills, house rules, ongoing communication, plus each roommate's \
personal notes.
- **Querier:** Either roommate (default: roommate A, designated in deep brief).
- **Cast:** Roommate A, Roommate B, landlord, building maintenance contact, \
occasional guests, neighbors mentioned in passing.
- **Directory shape:** `house/{rules,bills,maintenance}/`, \
`journal/{shared,personal_a,personal_b}/`, `messages/`, `memory/profiles/`.
- **File mix:** ~25 shared journal entries (some by A, some by B), house rules doc, \
monthly bills (Internet, utilities, rent split), Venmo logs, maintenance ticket \
emails, group chat exports, both roommates' personal scratch notes, profile of querier.
- **Eval stressors:** temporal recall ("when did the AC break?"), per-entity \
longitudinal across the two roommates, edit-then-recall, single-hop into specific \
bill files.
- **Realism notes:** small frictions (one tidies more, one cooks more), one shared \
amusement (running joke), realistic two-person banter.
""",
},
"dp_006": {
"file_count": 100,
"scenario_block": """\
## dp_006 — Indie open-source project, 6 months
- **Files:** 100
- **Setting:** A solo maintainer running a moderately popular open-source project \
(devtool, library, or CLI). 6 months of activity.
- **Querier:** The maintainer.
- **Cast:** The maintainer, ~10-20 community contributors (issue authors, PR \
submitters), 1-2 sponsors or notable users, occasional security-disclosure \
correspondent.
- **Directory shape:** `code/{rfcs,adr}/`, `issues/`, `pr_threads/`, `releases/`, \
`email/{users,sponsors,disclosure}/`, `notes/`, `memory/profiles/`.
- **File mix:** README, RFC and ADR docs, ~50 PR / issue threads (many authors, \
mixed lengths), 6 release-note files (one per month), changelog, sponsor outreach \
emails, security disclosure exchange (one), maintainer's scratch / planning notes, \
profile.
- **Eval stressors:** decision archaeology ("when did we drop Python 3.9 support \
and why?"), multi-hop (issue -> PR -> release note -> user follow-up), code-doc \
cross-references.
- **Realism notes:** GitHub-shaped issue/PR formats, authors with varying tone, \
drive-by issues, helpful regulars.
""",
},
"dp_007": {
"file_count": 200,
"scenario_block": """\
## dp_007 — Grad-student lab, first semester
- **Files:** 200
- **Setting:** A first-year PhD student's first semester. Lab is part of a larger \
department; advisor + 4 senior peers + a postdoc. Mix of coursework and lab work.
- **Querier:** The first-year PhD student.
- **Cast:** Student, advisor, postdoc, 4 lab peers, ~5 cohort classmates, professors \
of 4 courses, a couple paper authors with whom the student emailed.
- **Directory shape:** `papers/`, `lectures/{course1,course2,course3,course4}/`, \
`lab/{notebook,meetings,literature}/`, `meetings/{advisor,1on1}/`, `email/`, \
`memory/profiles/`.
- **File mix:** ~40 paper PDFs (with extracted-text sidecars), ~50 lecture notes \
(4 courses x ~12 weeks), ~25 problem sets and homework, lab notebook entries, weekly \
advisor 1:1 logs, group lab meeting notes, ~30 emails, departmental announcements, \
profile.
- **Eval stressors:** format-spanning (paper PDFs + sidecars are central), citation \
chains across reading + lecture notes, single-hop into specific lab notebook entries, \
temporal recall across the semester.
- **Realism notes:** real-shape academic correspondence; reading lists with \
annotations; "I should reread this" margin notes.
""",
},
"dp_008": {
"file_count": 300,
"scenario_block": """\
## dp_008 — Pre-seed startup, first 6 months
- **Files:** 300
- **Setting:** A 5-person pre-seed startup in its first 6 months. Two founders, \
2 early engineers, 1 designer/PM. Plus advisors and investors in correspondence.
- **Querier:** Any founder (default: CEO co-founder).
- **Cast:** 2 founders, 3 early team, 4-6 advisors, 8-12 investors / prospective \
investors, customer-interview subjects (10+), accountants, lawyers (incorporation), \
recruiter contact.
- **Directory shape:** `investors/{outreach,decks,follow_ups}/`, \
`customers/{interviews,demos}/`, `team/{slack,1on1}/`, `hiring/`, `legal/`, \
`decks/`, `memory/profiles/`.
- **File mix:** investor outreach emails, pitch deck iterations (3-4), customer \
interview transcripts (15+), co-founder slack export, hiring email threads, \
accounting/legal incorporation docs, founder's strategy memos, advisor emails, \
profile of querier.
- **Eval stressors:** profile.md heavy (founder context), multi-hop (advisor said \
X -> strategy memo references -> investor pitch reflects), founder-narrative \
coherence over time, edit-then-recall (note today's standup, retrieve next week).
- **Realism notes:** earnest scrappy energy, calendar friction, half-finished \
thoughts, optimism that pivots.
""",
},
"dp_009": {
"file_count": 500,
"scenario_block": """\
## dp_009 — Small therapy practice, 6 months, 4 therapists
- **Files:** 500
- **Setting:** A 4-therapist practice with a shared admin support, 6-month archive. \
Each therapist has 8-12 active clients, rotating; ~30 unique anonymized clients \
across the practice.
- **Querier:** Any therapist (default: senior therapist, 6-year licensed).
- **Cast:** 4 therapists, 1 admin/billing assistant, ~30 clients (anonymized as \
initials + ID), supervisor (external, monthly), insurance contacts.
- **Directory shape:** `clients/<client_id>/`, `staff/{notes,supervision}/`, \
`admin/{billing,scheduling,intake}/`, `ce_reading/`, `memory/profiles/`.
- **File mix:** ~360 anonymized session notes (12 clients x 30 sessions each on \
average per therapist; partitioned across the 4 therapists), supervisor session \
notes, CE reading notes, conference talk notes, intake forms, insurance \
correspondence, scheduling exports, profile of querier.
- **Eval stressors:** per-client longitudinal across many clients (distractor \
density), careful identity boundaries (do not leak across clients), ethical-shape \
correspondence.
- **Realism notes:** session notes follow SOAP-ish format; tone is careful and \
clinical; ethical boundaries explicit.
""",
},
"dp_010": {
"file_count": 1000,
"scenario_block": """\
## dp_010 — Growth-stage startup, 6 months, ~50 employees
- **Files:** 1,000
- **Setting:** A Series-A -> Series-B SaaS company, ~50 employees, 6-month archive. \
Multiple teams (eng, product, sales, CX, ops, exec). Multi-channel comms.
- **Querier:** A team lead (eng team lead by default; mid-level manager, ~20 reports \
including ICs and a coordinator).
- **Cast:** ~30-50 named employees (cross-team), 8-12 customers in active threads, \
2-3 vendors, 1 board member.
- **Directory shape:** `slack/{channel}/`, `email/{internal,customers,vendors}/`, \
`docs/{rfcs,post_mortems,playbooks}/`, `projects/{name}/`, `1on1/{report}/`, \
`meetings/{retros,planning,allhands}/`, `memory/profiles/`.
- **File mix:** Slack channel snapshots, email threads, design docs, post-mortems, \
project briefs, weekly 1:1 logs, retros, sprint planning notes, customer call notes, \
vendor correspondence, profile.
- **Eval stressors:** the "default" SMFS use case broadest test, all four task \
families, multi-hop across people/projects/time, profile.md heavy.
- **Realism notes:** cross-team noise, slack-shape banter, realistic name density.
""",
},
"dp_011": {
"file_count": 2000,
"scenario_block": """\
## dp_011 — Newsroom investigation, 18 months
- **Files:** 2,000
- **Setting:** A long-form investigative team (4 reporters + 2 editors) on a \
multi-month investigation. 18-month archive.
- **Querier:** Lead reporter.
- **Cast:** 4 reporters, 2 editors, ~20-30 sources (interviewed; varying anonymity), \
FOIA-respondent agencies, fact-checker, libel lawyer, photographer, a competing \
newsroom contact.
- **Directory shape:** `interviews/{audio,transcripts}/`, \
`sources/{notes,protected}/`, `foia/{requests,responses}/`, \
`editorial/{drafts,notes}/`, `published/`, `memory/profiles/`.
- **File mix:** ~50 interview audio transcripts (with sidecar text), source notes, \
FOIA correspondence + response PDFs (transcribed), editor email threads, draft \
article versions (multiple stages), background reading, fact-checking notes, photo \
logs (with image transcriptions), profile.
- **Eval stressors:** format-spanning is the headline (audio transcripts + FOIA \
PDFs), source-protection patterns, multi-hop across sources and documents.
- **Realism notes:** anonymized source IDs, sealed-source protocols, careful \
citation tracking.
""",
},
"dp_012": {
"file_count": 5000,
"scenario_block": """\
## dp_012 — Embassy at one posting, 3-year archive
- **Files:** 5,000
- **Setting:** A US embassy at one country posting, 3-year archive. Mid-size \
embassy: ~80 American staff, ~120 locally-employed staff, regular cable traffic.
- **Querier:** A mid-career FSO (Foreign Service Officer) on her second posting.
- **Cast:** US staff (DCM, political officer, econ officer, consular officers, RSO, \
defense attache, public affairs, Marine guard chief), locally-employed staff (LE \
staff key contacts), foreign-government counterparts (~30 named), regional NGO \
contacts, US-side desk officers in DC.
- **Directory shape:** `cables/{outgoing,incoming}/`, \
`briefings/{principals,vips}/`, `country/{political,economic,security}/`, \
`meetings/{readouts}/`, `personnel/`, `consular/`, `crisis/`, `memory/profiles/`.
- **File mix:** ~2,000 cables (varied classification levels reflected in metadata), \
briefing memos for visiting principals, country reports, meeting readouts, consular \
incident logs, crisis-response files, personnel-management notes, profile.
- **Eval stressors:** cross-relationship reasoning across foreign counterparts and \
DC desk officers, hierarchical surfaces (memos to ambassador, memos from DCM), \
classified-shape filing without inventing real classifications, temporal recall \
across postings.
- **Realism notes:** State-cable formatting (subjects, refs, drafted-by, cleared-by), \
realistic countries-fictional pairing (do not name a real geopolitical incident), \
tone is professional and indirect.
""",
},
"dp_013": {
"file_count": 10000,
"scenario_block": """\
## dp_013 — Tech-company CEO, full annual archive
- **Files:** 10,000
- **Setting:** A Series-B / early-Series-C tech company, ~300 employees, full \
12-month archive of the CEO's communications and accessible memory. Multiple \
departments, multiple ongoing projects, board, investors, customers, hiring, \
finance, HR.
- **Querier:** The CEO (or chief of staff acting on behalf).
- **Cast:** 200-500 named individuals: ~300 internal employees (sampled most \
active 100 appear repeatedly), ~30 board / investors, ~50 customers in active \
threads, ~20 vendors, ~30 candidate threads, family / personal life mixed in \
lightly, ~10 industry peers.
- **Directory shape:** \
`departments/{eng,product,sales,cx,ops,marketing,hr,finance,legal}/`, \
`projects/{name}/`, `slack/{channel}/`, `email/{internal,external,personal}/`, \
`board/{decks,minutes,prep}/`, `customers/{escalations,calls}/`, \
`hiring/{panels,decisions}/`, `finance/{reports,decisions}/`, \
`hr/{policies,sensitive}/`, `media/{interviews,press}/`, `memory/profiles/`.
- **File mix:** board decks and minutes, weekly 1:1 transcripts with 8 directs \
(x 52 weeks ~ 416), department-head reports, all-hands transcripts, hiring panel \
feedback, financial reports, HR matters (with care), customer escalation threads, \
strategy memos, investor communications, media interviews, industry-conference \
talks, daily executive-assistant briefings, personal email mixed in, profile.
- **Eval stressors:** highest-stakes profile.md (CEO context), multi-thread \
synthesis at scale, all four task families, edit-then-recall is high-stakes \
("Maya, write the board prep note for next quarter"), distractor robustness \
because of sheer corpus size.
- **Realism notes:** CEO voice is consistent across all CEO-authored files; varied \
tones across the named cast; realistic confidentiality boundaries on HR / financial \
files; strategy-decision arcs trace across multiple files.
""",
},
}
# ---------------------------------------------------------------------------
# Pipeline
# ---------------------------------------------------------------------------
async def run_pipeline(
dp_id: str,
*,
output_base: Path,
model: str = DEFAULT_MODEL,
max_concurrent: int = 10,
questions_only: bool = False,
validate_only: bool = False,
resume: bool = False,
) -> None:
"""Run the full generation pipeline for a single data point."""
scenario = SCENARIOS.get(dp_id)
if scenario is None:
raise ValueError(f"Unknown data point: {dp_id}. Available: {sorted(SCENARIOS.keys())}")
file_count = scenario["file_count"]
scenario_block = scenario["scenario_block"]
output_dir = output_base / dp_id
logger.info("=" * 60)
logger.info("Starting %s (%d files)", dp_id, file_count)
logger.info("Output: %s", output_dir)
logger.info("Model: %s", model)
logger.info("=" * 60)
# --- Validate-only mode ---
if validate_only:
manifest_path = output_dir / "manifest.json"
facts_path = output_dir / "facts.json"
if not manifest_path.exists() or not facts_path.exists():
logger.error("Cannot validate: manifest.json or facts.json missing in %s", output_dir)
return
manifest = read_json(manifest_path)
facts = read_json(facts_path)
report = await validate_corpus(output_dir, manifest, facts)
logger.info("Validation: %d errors, %d warnings out of %d files",
len(report.errors), len(report.warnings), report.total_files)
write_json(output_dir / "validation_report.json", {
"total_files": report.total_files,
"files_checked": report.files_checked,
"errors": len(report.errors),
"warnings": len(report.warnings),
"token_stats": report.token_stats,
"issues": [
{"file_id": i.file_id, "type": i.issue_type, "severity": i.severity,
"description": i.description}
for i in report.issues
],
})
return
# --- Questions-only mode ---
if questions_only:
scenario_path = output_dir / "SCENARIO.md"
facts_path = output_dir / "facts.json"
manifest_path = output_dir / "manifest.json"
if not all(p.exists() for p in [scenario_path, facts_path, manifest_path]):
logger.error("Cannot generate questions: missing SCENARIO.md, facts.json, or manifest.json")
return
brief = read_text(scenario_path)
facts = read_json(facts_path)
manifest = read_json(manifest_path)
questions = await generate_questions(output_dir, brief, facts, manifest, model=model)
logger.info("Generated %d questions -> %s/question.json", len(questions), output_dir)
return
# --- Full pipeline ---
t0 = time.monotonic()
# Phase 1-3: Planning
logger.info("--- PLANNING (Phases 1-3) ---")
brief, facts, manifest = await run_planning(
scenario_block=scenario_block,
file_count=file_count,
output_dir=output_dir,
model=model,
)
logger.info("Planning complete: %d facts categories, %d manifest entries",
len(facts), len(manifest))
# Phase 4: Clustering
logger.info("--- CLUSTERING (Phase 4) ---")
clusters = assign_clusters(manifest, facts)
logger.info("Created %d clusters across %d levels",
len(clusters), max(c.level for c in clusters) + 1 if clusters else 0)
# Phase 5: File generation
logger.info("--- GENERATION (Phase 5) ---")
gen_log = GenerationLog(output_dir / "generation_log.json")
# Build manifest lookup
manifest_entries = {e["file_id"]: e for e in manifest}
await generate_all(
clusters=clusters,
manifest_entries=manifest_entries,
output_dir=output_dir,
model=model,
max_concurrent=max_concurrent,
gen_log=gen_log,
fallback_fact_registry=facts,
)
gen_summary = gen_log.summary()
logger.info("Generation summary: %s", gen_summary)
# Phase 6: Validation
logger.info("--- VALIDATION (Phase 6) ---")
report = await validate_corpus(output_dir, manifest, facts)
logger.info("Validation: %d errors, %d warnings out of %d files",
len(report.errors), len(report.warnings), report.total_files)
write_json(output_dir / "validation_report.json", {
"total_files": report.total_files,
"files_checked": report.files_checked,
"errors": len(report.errors),
"warnings": len(report.warnings),
"token_stats": report.token_stats,
"issues": [
{"file_id": i.file_id, "type": i.issue_type, "severity": i.severity,
"description": i.description}
for i in report.issues
],
})
# Phase 7: Questions
logger.info("--- QUESTIONS (Phase 7) ---")
questions = await generate_questions(output_dir, brief, facts, manifest, model=model)
logger.info("Generated %d questions", len(questions))
elapsed = time.monotonic() - t0
logger.info("=" * 60)
logger.info("%s complete in %.1f minutes", dp_id, elapsed / 60)
logger.info(" Files: %d generated, %d failed",
gen_summary.get("ok", 0), gen_summary.get("failed", 0))
logger.info(" Tokens: %d in, %d out",
gen_summary.get("total_tokens_in", 0), gen_summary.get("total_tokens_out", 0))
logger.info(" Validation: %d errors, %d warnings",
len(report.errors), len(report.warnings))
logger.info(" Questions: %d", len(questions))
logger.info("=" * 60)
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Generate eval corpus data for memory benchmarks.",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""\
Examples:
python generate.py dp_001 # Generate dp_001
python generate.py dp_001 dp_005 # Generate dp_001 through dp_005
python generate.py dp_003 --resume # Resume failed dp_003
python generate.py dp_001 --questions-only # Only generate questions
python generate.py dp_002 --validate-only # Only validate existing corpus
""",
)
parser.add_argument(
"dp_start",
help="Data point ID to generate (e.g., dp_001)",
)
parser.add_argument(
"dp_end",
nargs="?",
default=None,
help="End of range (inclusive). If omitted, only dp_start is generated.",
)
parser.add_argument(
"--model",
default=DEFAULT_MODEL,
help=f"LLM model to use (default: {DEFAULT_MODEL})",
)
parser.add_argument(
"--max-concurrent",
type=int,
default=10,
help="Max concurrent cluster workers (default: 10)",
)
parser.add_argument(
"--output-dir",
type=Path,
default=Path("output"),
help="Base output directory (default: ./output)",
)
parser.add_argument(
"--questions-only",
action="store_true",
help="Only generate questions for an existing corpus",
)
parser.add_argument(
"--validate-only",
action="store_true",
help="Only validate an existing corpus",
)
parser.add_argument(
"--resume",
action="store_true",
help="Resume a failed generation (skip already-generated files)",
)
parser.add_argument(
"--log-level",
default="INFO",
choices=["DEBUG", "INFO", "WARNING", "ERROR"],
help="Logging level (default: INFO)",
)
return parser.parse_args()
def get_dp_range(start: str, end: str | None) -> list[str]:
"""Return list of dp IDs from start to end (inclusive)."""
all_dps = sorted(SCENARIOS.keys())
if start not in all_dps:
print(f"Error: unknown data point '{start}'. Available: {all_dps}")
sys.exit(1)
if end is None:
return [start]
if end not in all_dps:
print(f"Error: unknown data point '{end}'. Available: {all_dps}")
sys.exit(1)
start_idx = all_dps.index(start)
end_idx = all_dps.index(end)
if end_idx < start_idx:
print(f"Error: end '{end}' comes before start '{start}'")
sys.exit(1)
return all_dps[start_idx : end_idx + 1]
async def main() -> None:
args = parse_args()
# Configure logging
logging.basicConfig(
level=getattr(logging, args.log_level),
format="%(asctime)s %(levelname)-8s %(name)s%(message)s",
datefmt="%H:%M:%S",
)
dp_ids = get_dp_range(args.dp_start, args.dp_end)
logger.info("Will process: %s", ", ".join(dp_ids))
for dp_id in dp_ids:
try:
await run_pipeline(
dp_id,
output_base=args.output_dir,
model=args.model,
max_concurrent=args.max_concurrent,
questions_only=args.questions_only,
validate_only=args.validate_only,
resume=args.resume,
)
except Exception:
logger.exception("Failed to process %s", dp_id)
# Continue with next dp
continue
if __name__ == "__main__":
asyncio.run(main())

562
data-generator/planner.py Normal file
View file

@ -0,0 +1,562 @@
"""Planning module for the eval corpus data generator.
Handles three sequential phases:
Phase 1 Scenario Brief (SCENARIO.md)
Phase 2 Fact Registry (facts.json)
Phase 3 File Manifest (manifest.json)
"""
from __future__ import annotations
import json
import logging
from pathlib import Path
from typing import Any
from utils import (
DEFAULT_MODEL,
FAST_MODEL,
count_tokens,
llm_call,
llm_call_json,
read_text,
write_json,
write_text,
)
from prompts.scenario_brief import (
SCENARIO_BRIEF_SYSTEM,
format_scenario_brief_prompt,
)
from prompts.fact_registry import (
FACT_REGISTRY_SYSTEM,
format_fact_registry_prompt,
)
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
LARGE_CORPUS_THRESHOLD = 50
CHUNK_SIZE = 30
# ---------------------------------------------------------------------------
# Manifest prompt templates
# ---------------------------------------------------------------------------
MANIFEST_SYSTEM = """\
You are a corpus architect. Given a scenario brief and fact registry you must produce
a JSON manifest describing every file in the corpus. Each entry specifies exactly
what a downstream file-generator worker needs to produce that file.
Rules:
- file_id values are sequential: f001, f002,
- target_tokens [min, max] must both be in [5000, 10000]
- locked_facts lists reference fact IDs from the registry be exhaustive
- cross_references must be bidirectional (if f001 references f002, f002 references f001)
- cluster_hint groups related files (e.g. "legal", "medical_records", "travel")
- brief is 2-3 sentences describing the file's content
- authors is a list of person IDs from the fact registry
"""
MANIFEST_PROMPT = """\
## Task
Generate a file manifest (JSON array) for the corpus described below.
## Scenario Brief
{scenario_brief}
## Fact Registry
```json
{fact_registry}
```
## Requirements
Generate exactly {file_count} file entries as a JSON array. Each entry must have:
- **file_id**: sequential ID (f001, f002, )
- **path**: relative path under data/ (e.g. "data/emails/booking_confirmation.eml")
- **format**: document format (markdown_prose, email_thread, transcript, legal_contract, \
lab_report, slack_export, csv_data, json_structured, etc.)
- **authors**: list of person IDs from the fact registry
- **date**: ISO 8601 date (YYYY-MM-DD)
- **target_tokens**: [min, max] both within [5000, 10000]
- **locked_facts**: list of fact IDs from the registry that MUST appear in this file
- **cross_references**: list of other file_ids this file references or is referenced by
- **cluster_hint**: group name for related files
- **brief**: 2-3 sentence description of contents
- **tone**: formal/casual/clinical/technical/etc.
- **format_notes**: specific formatting requirements
Return ONLY a JSON array no wrapper object, no markdown fences.
"""
OUTLINE_PROMPT = """\
## Task
You are planning a large corpus of {file_count} files. To manage complexity, first
produce a department/section outline that organizes the files into logical groups.
## Scenario Brief (Summary)
{scenario_summary}
## Requirements
Return a JSON object with this structure:
```json
{{
"sections": [
{{
"name": "Section Name",
"cluster_hint": "section_slug",
"file_count": 15,
"description": "What files in this section cover"
}}
]
}}
```
Rules:
- Total file_count across all sections must equal exactly {file_count}
- Each section should have roughly {chunk_size} files (±10)
- Section names should be descriptive (e.g. "Legal Documents", "Medical Records")
- cluster_hint must be a URL-safe slug
"""
SECTION_MANIFEST_PROMPT = """\
## Task
Generate file manifest entries for the "{section_name}" section of the corpus.
## Scenario Brief
{scenario_brief}
## Fact Registry
```json
{fact_registry}
```
## Section Details
- **Section**: {section_name} ({section_description})
- **Cluster Hint**: {cluster_hint}
- **File Count**: {section_file_count}
- **Starting file_id**: f{start_id:03d}
## Requirements
Generate exactly {section_file_count} file entries as a JSON array. Each entry must have:
- **file_id**: sequential starting from f{start_id:03d}
- **path**: relative path under data/ (e.g. "data/{cluster_hint}/filename.ext")
- **format**: document format
- **authors**: list of person IDs from the fact registry
- **date**: ISO 8601 date (YYYY-MM-DD)
- **target_tokens**: [min, max] both within [5000, 10000]
- **locked_facts**: list of fact IDs from the registry that MUST appear in this file
- **cross_references**: list of other file_ids this file references (use IDs from any section)
- **cluster_hint**: "{cluster_hint}"
- **brief**: 2-3 sentence description of contents
- **tone**: formal/casual/clinical/technical/etc.
- **format_notes**: specific formatting requirements
Return ONLY a JSON array no wrapper object, no markdown fences.
"""
# ---------------------------------------------------------------------------
# Validation helpers
# ---------------------------------------------------------------------------
def _validate_fact_registry(registry: dict[str, Any]) -> dict[str, Any]:
"""Validate the fact registry has the expected top-level keys.
Returns the registry unchanged if valid, raises ValueError otherwise.
"""
required_keys = {"people", "organizations", "dates"}
missing = required_keys - set(registry.keys())
if missing:
raise ValueError(f"Fact registry missing required keys: {missing}")
# Validate people entries have 'id' fields
for person in registry.get("people", []):
if "id" not in person:
raise ValueError(f"Person entry missing 'id': {person}")
return registry
def _validate_manifest_entry(entry: dict[str, Any], idx: int) -> list[str]:
"""Validate a single manifest entry. Returns list of warnings (empty = OK)."""
warnings: list[str] = []
required_fields = {
"file_id", "path", "format", "authors", "date",
"target_tokens", "locked_facts", "cross_references",
"cluster_hint", "brief", "tone", "format_notes",
}
missing = required_fields - set(entry.keys())
if missing:
warnings.append(f"Entry {idx} missing fields: {missing}")
# Validate target_tokens range
tokens = entry.get("target_tokens")
if isinstance(tokens, list) and len(tokens) == 2:
lo, hi = tokens
if not (5000 <= lo <= 10000 and 5000 <= hi <= 10000):
warnings.append(
f"Entry {idx} target_tokens {tokens} outside [5000, 10000]"
)
if lo > hi:
warnings.append(f"Entry {idx} target_tokens min > max: {tokens}")
elif tokens is not None:
warnings.append(f"Entry {idx} target_tokens malformed: {tokens}")
return warnings
def _validate_manifest(manifest: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Validate all manifest entries. Logs warnings but returns the manifest."""
all_warnings: list[str] = []
for idx, entry in enumerate(manifest):
all_warnings.extend(_validate_manifest_entry(entry, idx))
if all_warnings:
for w in all_warnings:
logger.warning(f"Manifest validation: {w}")
return manifest
def _renumber_manifest(manifest: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Re-number file_ids sequentially (f001, f002, …) and update cross_references."""
# Build old-id → new-id mapping
id_map: dict[str, str] = {}
for idx, entry in enumerate(manifest):
old_id = entry.get("file_id", "")
new_id = f"f{idx + 1:03d}"
id_map[old_id] = new_id
entry["file_id"] = new_id
# Remap cross_references
for entry in manifest:
refs = entry.get("cross_references", [])
entry["cross_references"] = [
id_map.get(ref, ref) for ref in refs
]
return manifest
# ---------------------------------------------------------------------------
# Phase 1: Scenario Brief
# ---------------------------------------------------------------------------
async def generate_scenario_brief(
scenario_block: str,
file_count: int,
output_dir: Path,
model: str = DEFAULT_MODEL,
) -> str:
"""Phase 1: Generate SCENARIO.md. Returns the brief text."""
output_path = output_dir / "SCENARIO.md"
# Resume support: skip if already exists
if output_path.exists():
logger.info("Phase 1 skipped — SCENARIO.md already exists")
return read_text(output_path)
logger.info("Phase 1: Generating scenario brief …")
prompt = format_scenario_brief_prompt(scenario_block, file_count)
brief = await llm_call(
prompt,
model=model,
system=SCENARIO_BRIEF_SYSTEM,
max_tokens=16384,
)
write_text(output_path, brief)
logger.info(
"Phase 1 complete — SCENARIO.md written (%d tokens)", count_tokens(brief)
)
return brief
# ---------------------------------------------------------------------------
# Phase 2: Fact Registry
# ---------------------------------------------------------------------------
async def extract_fact_registry(
scenario_brief: str,
output_dir: Path,
model: str = DEFAULT_MODEL,
) -> dict:
"""Phase 2: Extract facts.json from SCENARIO.md. Returns the registry dict."""
output_path = output_dir / "facts.json"
# Resume support: skip if already exists
if output_path.exists():
logger.info("Phase 2 skipped — facts.json already exists")
data = json.loads(read_text(output_path))
return data
logger.info("Phase 2: Extracting fact registry …")
prompt = format_fact_registry_prompt(scenario_brief)
registry = await llm_call_json(
prompt,
model=model,
system=FACT_REGISTRY_SYSTEM,
max_tokens=16384,
)
# Handle case where registry is wrapped in a key
if isinstance(registry, dict) and len(registry) == 1:
key = next(iter(registry))
if isinstance(registry[key], dict):
# Might be double-wrapped; check if inner dict has expected keys
inner = registry[key]
if "people" in inner or "organizations" in inner:
registry = inner
_validate_fact_registry(registry)
write_json(output_path, registry)
fact_count = sum(
len(v) for v in registry.values() if isinstance(v, list)
)
logger.info("Phase 2 complete — facts.json written (%d fact entries)", fact_count)
return registry
# ---------------------------------------------------------------------------
# Phase 3: File Manifest
# ---------------------------------------------------------------------------
async def _generate_small_manifest(
scenario_brief: str,
fact_registry: dict,
file_count: int,
model: str,
) -> list[dict]:
"""Generate manifest in a single LLM call (≤50 files)."""
prompt = MANIFEST_PROMPT.format(
scenario_brief=scenario_brief,
fact_registry=json.dumps(fact_registry, indent=2),
file_count=file_count,
)
result = await llm_call_json(
prompt,
model=model,
system=MANIFEST_SYSTEM,
max_tokens=16384,
)
# Handle wrapped response — the LLM may return {"files": [...]} or similar
if isinstance(result, dict):
for key in ("files", "manifest", "entries"):
if key in result and isinstance(result[key], list):
return result[key]
# If it's a dict but no known key, look for any list value
for v in result.values():
if isinstance(v, list):
return v
raise ValueError(
f"Expected a JSON array for manifest, got dict with keys: {list(result.keys())}"
)
if isinstance(result, list):
return result
raise ValueError(f"Unexpected manifest response type: {type(result)}")
async def _generate_large_manifest(
scenario_brief: str,
fact_registry: dict,
file_count: int,
model: str,
) -> list[dict]:
"""Generate manifest in chunks for large corpora (>50 files)."""
# Summarize the brief if it's very long to keep section prompts under limit
brief_tokens = count_tokens(scenario_brief)
if brief_tokens > 6000:
scenario_summary = scenario_brief[:12000] + "\n\n[… truncated for outline …]"
else:
scenario_summary = scenario_brief
# Step 1: Generate section outline
logger.info("Phase 3a: Generating section outline for %d files …", file_count)
outline_prompt = OUTLINE_PROMPT.format(
file_count=file_count,
scenario_summary=scenario_summary,
chunk_size=CHUNK_SIZE,
)
outline = await llm_call_json(
outline_prompt,
model=model,
system=MANIFEST_SYSTEM,
max_tokens=4096,
)
sections = outline.get("sections", [])
if not sections:
raise ValueError("Outline generation returned no sections")
# Adjust section file counts to match total exactly
total_assigned = sum(s["file_count"] for s in sections)
if total_assigned != file_count:
diff = file_count - total_assigned
# Distribute difference across sections
sections[-1]["file_count"] += diff
logger.warning(
"Adjusted last section file_count by %d to match total %d",
diff,
file_count,
)
logger.info(
"Outline has %d sections: %s",
len(sections),
", ".join(f'{s["name"]}({s["file_count"]})' for s in sections),
)
# Step 2: Generate manifest for each section
all_entries: list[dict] = []
current_start_id = 1
for section in sections:
section_name = section["name"]
section_file_count = section["file_count"]
cluster_hint = section.get("cluster_hint", section_name.lower().replace(" ", "_"))
section_description = section.get("description", "")
logger.info(
"Phase 3b: Generating %d entries for section '%s' (starting f%03d) …",
section_file_count,
section_name,
current_start_id,
)
section_prompt = SECTION_MANIFEST_PROMPT.format(
section_name=section_name,
scenario_brief=scenario_brief,
fact_registry=json.dumps(fact_registry, indent=2),
section_description=section_description,
cluster_hint=cluster_hint,
section_file_count=section_file_count,
start_id=current_start_id,
)
result = await llm_call_json(
section_prompt,
model=model,
system=MANIFEST_SYSTEM,
max_tokens=16384,
)
# Extract list from potential wrapper
entries: list[dict]
if isinstance(result, list):
entries = result
elif isinstance(result, dict):
for key in ("files", "manifest", "entries"):
if key in result and isinstance(result[key], list):
entries = result[key]
break
else:
for v in result.values():
if isinstance(v, list):
entries = v
break
else:
raise ValueError(
f"Section '{section_name}' returned unexpected dict: "
f"{list(result.keys())}"
)
else:
raise ValueError(
f"Section '{section_name}' returned unexpected type: {type(result)}"
)
all_entries.extend(entries)
current_start_id += section_file_count
return all_entries
async def generate_manifest(
scenario_brief: str,
fact_registry: dict,
file_count: int,
output_dir: Path,
model: str = DEFAULT_MODEL,
) -> list[dict]:
"""Phase 3: Generate manifest.json. Returns list of file entries."""
output_path = output_dir / "manifest.json"
# Resume support: skip if already exists
if output_path.exists():
logger.info("Phase 3 skipped — manifest.json already exists")
data = json.loads(read_text(output_path))
if isinstance(data, dict) and "files" in data:
return data["files"]
return data
logger.info("Phase 3: Generating manifest for %d files …", file_count)
if file_count <= LARGE_CORPUS_THRESHOLD:
manifest = await _generate_small_manifest(
scenario_brief, fact_registry, file_count, model
)
else:
manifest = await _generate_large_manifest(
scenario_brief, fact_registry, file_count, model
)
# Re-number sequentially and fix cross-references
manifest = _renumber_manifest(manifest)
_validate_manifest(manifest)
write_json(output_path, manifest)
logger.info("Phase 3 complete — manifest.json written (%d entries)", len(manifest))
return manifest
# ---------------------------------------------------------------------------
# Orchestrator
# ---------------------------------------------------------------------------
async def run_planning(
scenario_block: str,
file_count: int,
output_dir: Path,
model: str = DEFAULT_MODEL,
) -> tuple[str, dict, list[dict]]:
"""Run all three planning phases sequentially.
Returns (brief, facts, manifest).
"""
out = Path(output_dir)
out.mkdir(parents=True, exist_ok=True)
brief = await generate_scenario_brief(scenario_block, file_count, out, model)
facts = await extract_fact_registry(brief, out, model)
manifest = await generate_manifest(brief, facts, file_count, out, model)
return brief, facts, manifest

View file

View file

@ -0,0 +1,120 @@
"""Prompt templates for Phase 2: Fact Registry extraction."""
FACT_REGISTRY_SYSTEM = """\
You are a precise data extractor. Your job is to read a scenario brief (SCENARIO.md)
and extract every concrete, verifiable fact into a structured JSON registry.
This registry is the single source of truth for corpus consistency. Every worker
generating files will receive this registry and must use these exact values.
Be exhaustive. If a fact appears in the brief, it must be in the registry.
If a fact could be referenced by multiple files, tag all of them.
"""
FACT_REGISTRY_PROMPT = """\
## Task
Extract a structured fact registry from the following scenario brief.
## Input: SCENARIO.md
{scenario_brief}
## Output: JSON
Return a JSON object with this exact structure:
```json
{{
"scenario_id": "dp_NNN",
"people": [
{{
"id": "person_slug",
"full_name": "Full Name",
"role": "Title, Organization",
"email": "email@example.com",
"timezone": "America/New_York",
"location": "City, State/Country",
"traits": ["trait1", "trait2"],
"writing_style": "Description of how they write",
"relationships": {{"person_slug2": "relationship description"}}
}}
],
"organizations": [
{{
"id": "org_slug",
"name": "Full Org Name",
"type": "company/hospital/embassy/etc",
"location": "City, State/Country",
"details": {{"key": "value"}}
}}
],
"dates": [
{{
"id": "date_slug",
"date": "YYYY-MM-DD",
"time": "HH:MM TZ (if applicable)",
"event": "What happened",
"participants": ["person_slug1", "person_slug2"],
"files": ["f001", "f002"]
}}
],
"financial": [
{{
"id": "financial_slug",
"value": "$X,XXX.XX",
"description": "What this amount represents",
"files": ["f001", "f003"]
}}
],
"references": [
{{
"id": "ref_slug",
"value": "EXACT-REF-CODE",
"type": "booking/case_number/mrn/confirmation/docket",
"description": "What this reference identifies",
"files": ["f002", "f004"]
}}
],
"locations": [
{{
"id": "location_slug",
"name": "Place Name",
"address": "Full address if known",
"type": "restaurant/hotel/office/hospital/etc",
"details": {{"key": "value"}},
"files": ["f001", "f005"]
}}
],
"domain_facts": [
{{
"id": "domain_slug",
"category": "medical/legal/technical/etc",
"fact": "The exact fact",
"files": ["f003", "f007"]
}}
],
"cross_references": [
{{
"source_file": "f001",
"target_file": "f002",
"fact_ids": ["financial_slug1", "date_slug2"],
"description": "How these files reference each other"
}}
]
}}
```
Rules:
- Every ID must be a unique, URL-safe slug (lowercase, underscores)
- Dollar amounts must include exact cents when relevant
- Dates must be ISO 8601 (YYYY-MM-DD)
- The "files" arrays must reference file_ids from the manifest (f001, f002, etc.)
- Include ALL facts, even minor ones completeness is critical
- Do not invent facts not in the brief
"""
def format_fact_registry_prompt(scenario_brief: str) -> str:
"""Format the fact registry extraction prompt."""
return FACT_REGISTRY_PROMPT.format(scenario_brief=scenario_brief)

View file

@ -0,0 +1,456 @@
"""Prompt templates for Phase 5: Individual file generation."""
FILE_GEN_SYSTEM = """\
You are generating a single realistic document for an eval corpus.
The document must feel like it was written by a real human in a real organization.
It is NOT a summary or a template it is the actual document itself.
Critical rules:
- Hit the target token count (5,000-10,000 tokens, roughly 20,000-40,000 characters)
- Include ALL locked facts exactly as specified these are the ground truth
- Match the specified format precisely (email headers for emails, speaker labels for transcripts, etc.)
- Write in the voice of the specified author(s)
- Include realistic noise: tangential discussions, filler, off-topic asides, pleasantries
- Do NOT be too organized or too clean real documents are messy
- Cross-references to other files should feel natural, not forced
- The document should be self-contained enough to read on its own, but clearly part of a larger corpus
"""
# ---------------------------------------------------------------------------
# Format-specific sub-prompts
# ---------------------------------------------------------------------------
FORMAT_INSTRUCTIONS: dict[str, str] = {
"email_thread": """\
Format: Email Thread
- Every message MUST have headers: From, To, Cc (if applicable), Date, Subject
- Reply subjects use "Re: ..." prefix
- Include realistic email signatures (name, title, phone, confidentiality disclaimers)
- Thread is ordered chronologically (oldest first)
- Include realistic forwarding artifacts ("---------- Forwarded message ----------")
- Vary reply lengths some are one-liners, some are paragraphs
- Include the occasional top-posted reply with full quote chain below
""",
"transcript": """\
Format: Meeting/Call Transcript
- Every utterance starts with a speaker label and timestamp: "[HH:MM:SS] Speaker Name:"
- Include filler words naturally: um, uh, like, you know, I mean, so, right
- Include crosstalk markers: [crosstalk], [overlapping], [inaudible]
- Include non-verbal cues: [laughs], [sighs], [pause], [typing sounds]
- Some speakers interrupt others mid-sentence
- Include an opening attendance/roll call section
- Include off-topic small talk at the beginning and end
""",
"legal_contract": """\
Format: Legal Contract / Agreement
- Start with a title block: agreement type, date, parties
- Include WHEREAS recital clauses
- Use numbered sections (1., 1.1, 1.1.1) with descriptive headings
- Include a Definitions section near the top
- Use defined terms in Title Case or ALL CAPS with quotes on first use
- Include standard boilerplate: governing law, severability, entire agreement, counterparts
- End with signature blocks (name, title, date lines)
- Use formal legal prose passive voice, shall/may distinctions
""",
"slack_export": """\
Format: Slack Channel Export
- Each message has a timestamp and username: "[YYYY-MM-DD HH:MM] @username:"
- Include thread replies indented or marked: " ↳ [HH:MM] @username:"
- Include emoji reactions: " :thumbsup: (3) :eyes: (1)"
- Include @mentions, #channel-references, and :emoji: usage
- Include bot messages (e.g., "/remind", integration notifications)
- Some messages are very short ("lol", "^", "+1", "👍")
- Include edited markers: "(edited)"
- Include file/link sharing: "[shared a file: quarterly_report.pdf]"
""",
"clinical_note": """\
Format: Clinical / Medical Note
- Follow SOAP format: Subjective, Objective, Assessment, Plan
- Include patient header: MRN, DOB, encounter date, provider, clinic
- Use standard medical abbreviations: pt, hx, dx, tx, prn, bid, tid, qd, etc.
- Include vitals in structured format: BP, HR, RR, Temp, SpO2, Weight
- Include medication lists with dosages and frequencies
- Use ICD-10 codes where appropriate
- Include review of systems (ROS) section
- Assessment uses clinical reasoning language
- Plan includes numbered action items
""",
"memo": """\
Format: Internal Memo / Memorandum
- Start with a header block: TO, FROM, DATE, RE (or SUBJECT)
- Use formal prose paragraphs
- May include numbered points or bullet lists for action items
- End with a signature or initials
- Tone is professional but can vary by organization culture
- May include "cc:" line at the bottom
""",
"markdown_prose": """\
Format: Markdown Document
- Use natural markdown formatting: #/##/### headers, **bold**, *italic*
- Include bullet lists, numbered lists, and occasional tables
- Use code blocks if technical content is relevant
- Include hyperlinks (can be realistic URLs or internal wiki links)
- Structure should feel like a real wiki page, report, or documentation
- Can include a table of contents for longer documents
""",
"profile": """\
Format: Personal / Professional Profile
- Include structured sections: name, role, contact info, bio
- Include professional background and expertise areas
- Include personal details relevant to the scenario (preferences, restrictions, etc.)
- Can be formatted as markdown, YAML-style, or structured text
- Include relevant metadata: timezone, location, team membership
- Feels like an HR system profile or internal directory entry
""",
}
# ---------------------------------------------------------------------------
# Anti-pattern warnings included in every prompt
# ---------------------------------------------------------------------------
ANTI_PATTERN_WARNINGS = """\
## Anti-Pattern Warnings — READ CAREFULLY
You are an AI generating this document. You MUST fight your instincts to be too clean:
- Do NOT use perfect grammar in casual communications (emails, Slack)
- Do NOT make every paragraph equally sized
- Do NOT include a perfect topic sentence for every section
- Do NOT organize information in a neat, logical order real documents ramble
- Do NOT make facts easy to find bury some in the middle of unrelated paragraphs
- Do NOT use headers or bullets where the real format wouldn't have them
- Do NOT summarize or conclude unless the format calls for it
- DO include tangential asides, personal anecdotes, and off-topic filler
- DO vary sentence length dramatically
- DO include some redundancy (same point made slightly differently)
- DO include realistic noise that adds length without adding information
"""
# ---------------------------------------------------------------------------
# Main prompt template
# ---------------------------------------------------------------------------
FILE_GEN_PROMPT = """\
## Task
Generate the complete content of the following document. Output ONLY the document \
content no wrapper, no explanation, no metadata.
## File Brief
- **File ID**: {file_id}
- **Path**: {file_path}
- **Format**: {file_format}
- **Date**: {file_date}
- **Author(s)**: {file_authors}
- **Tone**: {file_tone}
- **Summary**: {file_summary}
## Target Length
{target_length_instructions}
This is critical. The document MUST be long enough. Pad with realistic filler, \
tangential discussion, pleasantries, and noise if needed. A document that is too \
short is a failure.
## Format-Specific Instructions
{format_instructions}
{format_notes}
## Author Information & Writing Style
{author_info}
## Locked Facts — MUST Appear in This Document
The following facts are ground truth. They MUST appear in the generated document \
exactly as specified. Do not alter names, numbers, dates, or reference codes.
{locked_facts}
## Cross-Reference Context
The following are other documents in the corpus that this file references or is \
related to. Use them for context, but do not copy them verbatim. References should \
feel natural.
{cross_reference_context}
{anti_pattern_warnings}
## Final Reminder
- Target: {target_min_chars}-{target_max_chars} characters ({target_min_tokens}-{target_max_tokens} tokens)
- Include ALL locked facts
- Write as the specified author(s), not as an AI
- The document should feel real, messy, and human
- Output ONLY the document content nothing else
"""
RETRY_FEEDBACK_PROMPT = """\
## Retry: Fix the Following Issues
Your previous attempt had these problems:
{issues}
## Previous Attempt (for reference)
{previous_attempt_truncated}
## Original Instructions
{original_prompt}
Please regenerate the COMPLETE document, fixing all listed issues. \
Output ONLY the document content.
"""
def _build_author_info(file_entry: dict, fact_shard: dict) -> str:
"""Build author information and writing style section from the fact shard."""
authors = file_entry.get("authors") or file_entry.get("author", [])
if isinstance(authors, str):
authors = [authors]
people = {p["id"]: p for p in fact_shard.get("people", [])}
parts: list[str] = []
for author_id in authors:
person = people.get(author_id)
if person:
lines = [
f"**{person.get('full_name', author_id)}**",
f"- Role: {person.get('role', 'Unknown')}",
f"- Email: {person.get('email', 'N/A')}",
f"- Location: {person.get('location', 'N/A')}",
f"- Timezone: {person.get('timezone', 'N/A')}",
]
if person.get("writing_style"):
lines.append(f"- Writing style: {person['writing_style']}")
if person.get("traits"):
lines.append(f"- Traits: {', '.join(person['traits'])}")
if person.get("relationships"):
rels = "; ".join(
f"{k}: {v}" for k, v in person["relationships"].items()
)
lines.append(f"- Relationships: {rels}")
parts.append("\n".join(lines))
else:
parts.append(f"**{author_id}** (no detailed profile available)")
return "\n\n".join(parts) if parts else "No author information available."
def _build_locked_facts(file_entry: dict, fact_shard: dict) -> str:
"""Build the locked facts section for a file."""
locked_ids = set(file_entry.get("locked_facts", []))
if not locked_ids:
return "No specific locked facts for this file."
fact_lines: list[str] = []
# Search across all fact categories
for category in [
"financial",
"dates",
"references",
"locations",
"domain_facts",
]:
for fact in fact_shard.get(category, []):
if fact.get("id") in locked_ids:
if category == "financial":
fact_lines.append(
f"- **[{category}]** {fact['id']}: "
f"{fact.get('value', '')}{fact.get('description', '')}"
)
elif category == "dates":
date_str = fact.get("date", "")
time_str = fact.get("time", "")
fact_lines.append(
f"- **[{category}]** {fact['id']}: "
f"{date_str} {time_str}{fact.get('event', '')}"
)
elif category == "references":
fact_lines.append(
f"- **[{category}]** {fact['id']}: "
f"{fact.get('value', '')} ({fact.get('type', '')}) — "
f"{fact.get('description', '')}"
)
elif category == "locations":
fact_lines.append(
f"- **[{category}]** {fact['id']}: "
f"{fact.get('name', '')}{fact.get('address', '')} "
f"({fact.get('type', '')})"
)
elif category == "domain_facts":
fact_lines.append(
f"- **[{category}]** {fact['id']}: {fact.get('fact', '')}"
)
# Also check people facts that might be locked
for person in fact_shard.get("people", []):
if person.get("id") in locked_ids:
fact_lines.append(
f"- **[person]** {person['id']}: "
f"{person.get('full_name', '')}{person.get('role', '')}"
)
# Also check organizations
for org in fact_shard.get("organizations", []):
if org.get("id") in locked_ids:
fact_lines.append(
f"- **[organization]** {org['id']}: "
f"{org.get('name', '')} ({org.get('type', '')})"
)
if not fact_lines:
return (
f"Locked fact IDs: {', '.join(sorted(locked_ids))}\n"
"(Could not resolve full details — use the IDs as-is from context.)"
)
return "\n".join(fact_lines)
def _build_cross_reference_context(
file_entry: dict, context_files: dict[str, str], manifest_entries: dict[str, dict]
) -> str:
"""Build the cross-reference context section.
Args:
file_entry: the manifest entry for the file being generated.
context_files: file_id -> content of already-generated files.
manifest_entries: file_id -> manifest entry for all files (for briefs).
"""
cross_refs = file_entry.get("cross_references", [])
if not cross_refs:
return "No cross-references for this file."
parts: list[str] = []
for ref_id in cross_refs:
if ref_id in context_files:
parts.append(
f"### {ref_id} (generated)\n\n{context_files[ref_id]}"
)
elif ref_id in manifest_entries:
entry = manifest_entries[ref_id]
brief = (
f"**{ref_id}** — {entry.get('path', 'unknown path')}\n"
f"Format: {entry.get('format', 'unknown')}\n"
f"Summary: {entry.get('summary', 'No summary available')}"
)
parts.append(f"### {ref_id} (not yet generated — brief only)\n\n{brief}")
else:
parts.append(f"### {ref_id}\n\n(No information available)")
return "\n\n---\n\n".join(parts) if parts else "No cross-references for this file."
def format_file_gen_prompt(
file_entry: dict,
fact_shard: dict,
context_files: dict[str, str],
manifest_entries: dict[str, dict] | None = None,
) -> tuple[str, str]:
"""Format the file generation prompt.
Args:
file_entry: manifest entry for this file.
fact_shard: relevant portion of fact registry.
context_files: file_id -> content for already-generated referenced files.
manifest_entries: file_id -> manifest entry for all files (used for
cross-reference briefs of not-yet-generated files).
Returns:
(system_prompt, user_prompt) tuple.
"""
if manifest_entries is None:
manifest_entries = {}
file_format = file_entry.get("format", "markdown_prose")
format_instructions = FORMAT_INSTRUCTIONS.get(
file_format, FORMAT_INSTRUCTIONS["markdown_prose"]
)
format_notes = file_entry.get("format_notes", "")
if format_notes:
format_notes = f"### Additional Format Notes\n\n{format_notes}"
# Compute target character counts from token range
target_tokens = file_entry.get("target_tokens", [5000, 10000])
target_min_tokens = target_tokens[0] if isinstance(target_tokens, list) else 5000
target_max_tokens = target_tokens[1] if isinstance(target_tokens, list) else 10000
target_min_chars = target_min_tokens * 4
target_max_chars = target_max_tokens * 4
target_length_instructions = (
f"- Target: **{target_min_tokens:,}-{target_max_tokens:,} tokens** "
f"(approximately **{target_min_chars:,}-{target_max_chars:,} characters**)\n"
f"- This means the document should be LONG. Think 5-10 pages of text.\n"
f"- Err on the side of MORE content, not less."
)
# Build authors string
authors = file_entry.get("authors") or file_entry.get("author", [])
if isinstance(authors, list):
authors_str = ", ".join(str(a) for a in authors)
else:
authors_str = str(authors)
prompt = FILE_GEN_PROMPT.format(
file_id=file_entry.get("file_id", "unknown"),
file_path=file_entry.get("path", "unknown"),
file_format=file_format,
file_date=file_entry.get("date", "unknown"),
file_authors=authors_str,
file_tone=file_entry.get("tone", "neutral"),
file_summary=file_entry.get("summary", "No summary provided."),
target_length_instructions=target_length_instructions,
format_instructions=format_instructions,
format_notes=format_notes,
author_info=_build_author_info(file_entry, fact_shard),
locked_facts=_build_locked_facts(file_entry, fact_shard),
cross_reference_context=_build_cross_reference_context(
file_entry, context_files, manifest_entries
),
anti_pattern_warnings=ANTI_PATTERN_WARNINGS,
target_min_chars=f"{target_min_chars:,}",
target_max_chars=f"{target_max_chars:,}",
target_min_tokens=f"{target_min_tokens:,}",
target_max_tokens=f"{target_max_tokens:,}",
)
return FILE_GEN_SYSTEM, prompt
def format_retry_prompt(
issues: list[str],
previous_content: str,
original_prompt: str,
max_previous_chars: int = 8000,
) -> str:
"""Format a retry prompt with feedback about what went wrong.
Args:
issues: list of issue descriptions.
previous_content: the previous attempt's content (will be truncated).
original_prompt: the original user prompt.
max_previous_chars: max characters to include from previous attempt.
Returns:
The formatted retry prompt.
"""
truncated = previous_content[:max_previous_chars]
if len(previous_content) > max_previous_chars:
truncated += "\n\n[... truncated ...]"
issues_str = "\n".join(f"- {issue}" for issue in issues)
return RETRY_FEEDBACK_PROMPT.format(
issues=issues_str,
previous_attempt_truncated=truncated,
original_prompt=original_prompt,
)

View file

@ -0,0 +1,54 @@
"""Prompt templates for Phase 7: Eval Question Generation."""
QUESTION_GEN_SYSTEM = """You are generating eval questions for a memory retrieval benchmark.
Each question tests whether a system can find and synthesize information from a corpus of files.
Question families:
- single_hop: answer is in one file, straightforward retrieval
- multi_hop: answer requires combining info from 2-3 files
- format_spanning: answer requires info from files in different formats (e.g., email + contract)
- edit_then_recall: answer involves a fact that was updated/changed across files
"""
QUESTION_GEN_PROMPT = """Generate exactly 10 eval questions for this corpus.
## Scenario
{scenario_summary}
## Fact Registry (key facts)
{fact_summary}
## File Manifest
{manifest_summary}
## Distribution
Generate approximately:
- 3 single_hop questions
- 3 multi_hop questions
- 2 format_spanning questions
- 2 edit_then_recall questions
## Output Format
Return a JSON array:
[
{{
"id": "q01",
"family": "single_hop",
"prompt": "The natural-language question an agent would be asked",
"gold_file_ids": ["data/path/to/file1.md", "data/path/to/file2.eml"],
"gold_answer": "The exact answer string"
}},
...
]
Rules:
- gold_file_ids are the file paths (relative to the dp directory) needed to answer
- gold_answer is the literal answer concise, factual, no hedging
- Questions should feel natural, like a real user asking their AI assistant
- single_hop questions should be answerable from exactly 1 file
- multi_hop questions should require 2-3 files
- format_spanning questions should require files of different formats
- edit_then_recall questions should involve facts that appear differently in different files
- Do NOT ask questions whose answers aren't in the corpus
- Do NOT ask meta-questions about the corpus itself
"""

View file

@ -0,0 +1,134 @@
"""Prompt templates for Phase 1: Scenario Brief generation."""
SCENARIO_BRIEF_SYSTEM = """\
You are a world-builder for synthetic eval corpora. Your job is to take a high-level
scenario description and produce a detailed "bible" the complete ground truth for
an entire organizational corpus.
The corpus simulates a real organization's shared memory: files written by many authors,
in many formats, over a specific time period. Every fact you establish becomes canonical.
Downstream workers will generate individual files from your brief, so you must be
exhaustive and precise.
Key principles:
- Every person has a distinct voice, background, and role
- Dates, dollar amounts, reference numbers, and proper nouns are LOCKED they must be
exact and consistent
- The corpus must feel like it was written by real humans, not AI
- Include realistic messiness: typos in casual messages, formal tone in contracts,
medical jargon in clinical notes, etc.
- Cross-references between files must be explicit and bidirectional
"""
SCENARIO_BRIEF_PROMPT = """\
## Task
Generate a comprehensive SCENARIO.md brief for the following eval data point.
## Input: Scenario Description
{scenario_block}
## Output Requirements
Produce a detailed markdown document with these sections:
### 1. Overview
- Scenario ID, file count, time span, setting
- One-paragraph narrative summary
### 2. Cast of Characters
For EVERY named person:
- Full name, role/title, organization
- Email address (realistic format)
- Timezone, location
- 2-3 personality/writing-style notes (e.g., "writes terse emails", "uses emoji in Slack")
- Key traits relevant to the scenario (dietary restrictions, allergies, expertise areas)
- Relationships to other cast members
### 3. Organizations
For every org/company/institution:
- Full name, type, location
- Key facts (size, industry, founding date if relevant)
- Internal structure relevant to the scenario
### 4. Timeline
A chronological list of every event in the scenario:
- Date (YYYY-MM-DD) and time if relevant
- What happened
- Who was involved
- Which files document this event
### 5. Locked Facts Registry
Every concrete fact that MUST be consistent across files. Group by category:
- **Financial**: dollar amounts, rates, costs, budgets
- **References**: booking refs, case numbers, docket numbers, MRNs, confirmation codes
- **Dates**: deadlines, appointments, milestones
- **Locations**: addresses, room numbers, restaurant names
- **Technical**: system names, version numbers, tool names
- **Medical/Legal/Domain**: diagnoses, statutes, specifications
Each fact must specify: the exact value, which files it appears in, and any context.
### 6. Directory Structure
The exact file tree for the corpus:
```
data/
[domain folders]/
file1.md
file2.eml
memory/
profiles/
[other memory subdirs]/
```
### 7. File Manifest
For EVERY file in the corpus, provide:
- **file_id**: f001, f002, ...
- **path**: relative path under data/
- **format**: the document format (markdown_prose, email_thread, transcript, legal_contract,
lab_report, slack_export, csv_data, json_structured, etc.)
- **author(s)**: who wrote/sent this
- **date**: when this was created/sent
- **target_tokens**: [min, max] within [5000, 10000]
- **summary**: 2-3 sentence description of what this file contains
- **locked_facts**: list of fact IDs from the registry that MUST appear in this file
- **cross_references**: list of other file_ids this file references or is referenced by
- **tone**: formal/casual/clinical/technical/etc.
- **format_notes**: specific formatting requirements (email headers, transcript speaker
labels, legal clause numbering, etc.)
### 8. Cross-Reference Map
A table showing every cross-reference between files:
| Source File | Target File | What's Referenced | Direction |
|-------------|-------------|-------------------|-----------|
### 9. Eval Stressor Notes
Which eval stressors this scenario tests and how:
- Single-hop retrieval targets
- Multi-hop chains (file A file B file C)
- Format-spanning queries (answer requires info from different file formats)
- Edit-then-recall patterns
- Profile/cheap-read targets
### 10. Anti-Pattern Warnings
Specific instructions for file generators to avoid AI-perfection:
- Which files should have typos or informal language
- Where noise/filler content should appear
- Which files should have near-zero extractable facts
- Where information should be buried rather than prominent
## Important
- Generate {file_count} files total, no more, no less
- Every file must target 5,000-10,000 tokens
- The memory/ directory must contain the querier's profile and relevant reference docs
- Do NOT reuse names from other scenarios (cross-scenario isolation)
- Be exhaustive downstream workers will generate files from this brief alone
"""
def format_scenario_brief_prompt(scenario_block: str, file_count: int) -> str:
"""Format the scenario brief prompt with the scenario description."""
return SCENARIO_BRIEF_PROMPT.format(
scenario_block=scenario_block,
file_count=file_count,
)

323
data-generator/questions.py Normal file
View file

@ -0,0 +1,323 @@
"""Phase 7: Eval Question Generation.
Generates eval questions for a corpus, testing retrieval across four families:
single_hop, multi_hop, format_spanning, and edit_then_recall.
"""
from __future__ import annotations
import json
import logging
from pathlib import Path
from typing import Any
from utils import DEFAULT_MODEL, count_tokens, llm_call_json, read_text, write_json
from prompts.questions import QUESTION_GEN_PROMPT, QUESTION_GEN_SYSTEM
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
SMALL_CORPUS_THRESHOLD = 50
MAX_SCENARIO_TOKENS = 2000
MAX_EXCERPT_TOKENS = 500
MAX_TOTAL_EXCERPT_TOKENS = 3000
VALID_FAMILIES = {"single_hop", "multi_hop", "format_spanning", "edit_then_recall"}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _truncate_to_tokens(text: str, max_tokens: int) -> str:
"""Truncate text to approximately max_tokens.
Uses a rough 4-chars-per-token estimate for speed, then verifies.
"""
if count_tokens(text) <= max_tokens:
return text
# Rough cut, then refine
char_estimate = max_tokens * 4
truncated = text[:char_estimate]
# Trim to last complete sentence or paragraph
for sep in ("\n\n", "\n", ". ", " "):
idx = truncated.rfind(sep)
if idx > char_estimate // 2:
truncated = truncated[: idx + len(sep)]
break
return truncated + "\n\n[… truncated …]"
def _build_scenario_summary(scenario_brief: str) -> str:
"""Build a truncated scenario summary for the prompt."""
return _truncate_to_tokens(scenario_brief, MAX_SCENARIO_TOKENS)
def _build_fact_summary(fact_registry: dict) -> str:
"""Build a concise summary of key facts from the registry."""
lines: list[str] = []
# People — just names and roles
people = fact_registry.get("people", [])
if people:
lines.append("### People")
for p in people:
name = p.get("full_name", p.get("id", "unknown"))
role = p.get("role", "")
lines.append(f"- {name}: {role}")
# Financial facts
financial = fact_registry.get("financial", [])
if financial:
lines.append("\n### Financial")
for f in financial:
lines.append(f"- {f.get('id', '')}: {f.get('value', '')}{f.get('description', '')}")
# References
references = fact_registry.get("references", [])
if references:
lines.append("\n### References")
for r in references:
lines.append(f"- {r.get('id', '')}: {r.get('value', '')} ({r.get('type', '')})")
# Key dates
dates = fact_registry.get("dates", [])
if dates:
lines.append("\n### Key Dates")
for d in dates:
lines.append(f"- {d.get('id', '')}: {d.get('date', '')}{d.get('event', '')}")
# Domain facts (first 10 only to stay concise)
domain = fact_registry.get("domain_facts", [])
if domain:
lines.append("\n### Domain Facts")
for df in domain[:10]:
lines.append(f"- {df.get('id', '')}: {df.get('fact', '')[:100]}")
if len(domain) > 10:
lines.append(f" … and {len(domain) - 10} more")
# Cross-references summary
xrefs = fact_registry.get("cross_references", [])
if xrefs:
lines.append(f"\n### Cross-References: {len(xrefs)} connections between files")
return "\n".join(lines)
def _build_manifest_summary(manifest: list[dict]) -> str:
"""Build a concise manifest summary showing file briefs and formats."""
lines: list[str] = []
for entry in manifest:
file_id = entry.get("file_id", "?")
path = entry.get("path", "?")
fmt = entry.get("format", "?")
brief = entry.get("brief", "")
date = entry.get("date", "")
locked = entry.get("locked_facts", [])
line = f"- **{file_id}** `{path}` ({fmt}, {date})"
if brief:
line += f": {brief[:120]}"
if locked:
line += f" [facts: {', '.join(locked[:5])}{'' if len(locked) > 5 else ''}]"
lines.append(line)
return "\n".join(lines)
def _sample_file_excerpts(
output_dir: Path,
manifest: list[dict],
) -> str:
"""Sample excerpts from a few files to ground questions in actual content.
Only used for small corpora (<=50 files).
"""
excerpts: list[str] = []
total_tokens = 0
# Sample up to 8 files, evenly distributed across the manifest
sample_count = min(8, len(manifest))
if sample_count == 0:
return ""
step = max(1, len(manifest) // sample_count)
sampled_entries = manifest[::step][:sample_count]
for entry in sampled_entries:
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
if not full_path.exists():
continue
content = read_text(full_path)
excerpt = _truncate_to_tokens(content, MAX_EXCERPT_TOKENS)
excerpt_tokens = count_tokens(excerpt)
if total_tokens + excerpt_tokens > MAX_TOTAL_EXCERPT_TOKENS:
break
file_id = entry.get("file_id", "?")
excerpts.append(f"### {file_id} ({rel_path})\n{excerpt}")
total_tokens += excerpt_tokens
if not excerpts:
return ""
return "\n\n---\n\n## Sample File Excerpts\n\n" + "\n\n".join(excerpts)
def _validate_questions(questions: list[dict], manifest: list[dict]) -> list[dict]:
"""Validate and clean up generated questions."""
valid_paths = {entry.get("path", "") for entry in manifest}
valid_file_ids = {entry.get("file_id", "") for entry in manifest}
# Build a file_id -> path mapping for normalization
id_to_path: dict[str, str] = {}
for entry in manifest:
fid = entry.get("file_id", "")
path = entry.get("path", "")
if fid and path:
id_to_path[fid] = path
validated: list[dict] = []
for q in questions:
# Ensure required fields
if not all(k in q for k in ("id", "family", "prompt", "gold_file_ids", "gold_answer")):
logger.warning("Skipping question missing required fields: %s", q.get("id", "?"))
continue
# Normalize family
family = q.get("family", "").lower().replace("-", "_")
if family not in VALID_FAMILIES:
logger.warning(
"Question %s has unknown family '%s', keeping as-is",
q.get("id", "?"),
family,
)
q["family"] = family
# Normalize gold_file_ids: convert file_ids to paths if needed
normalized_ids: list[str] = []
for gid in q.get("gold_file_ids", []):
if gid in valid_paths:
normalized_ids.append(gid)
elif gid in id_to_path:
normalized_ids.append(id_to_path[gid])
elif gid in valid_file_ids:
# It's a valid file_id but has no path mapping (shouldn't happen)
normalized_ids.append(gid)
else:
logger.warning(
"Question %s references unknown file '%s'",
q.get("id", "?"),
gid,
)
normalized_ids.append(gid)
q["gold_file_ids"] = normalized_ids
validated.append(q)
return validated
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
async def generate_questions(
output_dir: Path,
scenario_brief: str,
fact_registry: dict,
manifest: list[dict],
model: str = DEFAULT_MODEL,
) -> list[dict]:
"""Generate 10 eval questions for a corpus.
Reads the generated files to understand what's actually in the corpus,
then generates questions that test retrieval across the four families.
Returns list of question dicts and writes to output_dir/question.json.
"""
output_path = output_dir / "question.json"
# Resume support: skip if already exists
if output_path.exists():
logger.info("Phase 7 skipped — question.json already exists")
return read_text(output_path)
logger.info("Phase 7: Generating eval questions …")
# Build summaries for the prompt (don't pass full file contents — too large)
scenario_summary = _build_scenario_summary(scenario_brief)
fact_summary = _build_fact_summary(fact_registry)
manifest_summary = _build_manifest_summary(manifest)
# For small corpora, also sample file excerpts to ground the questions
excerpt_section = ""
if len(manifest) <= SMALL_CORPUS_THRESHOLD:
excerpt_section = _sample_file_excerpts(output_dir, manifest)
prompt = QUESTION_GEN_PROMPT.format(
scenario_summary=scenario_summary,
fact_summary=fact_summary,
manifest_summary=manifest_summary,
)
# Append excerpts if available
if excerpt_section:
prompt += excerpt_section
result = await llm_call_json(
prompt,
model=model,
system=QUESTION_GEN_SYSTEM,
max_tokens=8192,
)
# Handle wrapped responses — LLM may return {"questions": [...]}
questions: list[dict]
if isinstance(result, list):
questions = result
elif isinstance(result, dict):
for key in ("questions", "eval_questions", "items"):
if key in result and isinstance(result[key], list):
questions = result[key]
break
else:
# Look for any list value
for v in result.values():
if isinstance(v, list):
questions = v
break
else:
raise ValueError(
f"Expected a JSON array for questions, got dict with keys: {list(result.keys())}"
)
else:
raise ValueError(f"Unexpected questions response type: {type(result)}")
# Validate and clean up
questions = _validate_questions(questions, manifest)
# Log distribution
family_counts: dict[str, int] = {}
for q in questions:
fam = q.get("family", "unknown")
family_counts[fam] = family_counts.get(fam, 0) + 1
logger.info(
"Phase 7 complete — %d questions generated: %s",
len(questions),
", ".join(f"{k}={v}" for k, v in sorted(family_counts.items())),
)
write_json(output_path, questions)
return questions

View file

@ -0,0 +1,6 @@
litellm>=1.50.0
pydantic>=2.0
pydantic-settings>=2.0
tiktoken>=0.7.0
pyyaml>=6.0
asyncio-pool>=0.7.0

View file

@ -0,0 +1,603 @@
"""Tests for clusterer.py — Phase 4: cluster assignment, topo sort, fact sharding."""
from __future__ import annotations
import warnings
import pytest
from clusterer import (
Cluster,
_build_cluster_dependency_graph,
_build_cross_ref_graph,
_detect_and_merge_cycles,
_find_connected_components,
_get_cluster_hint,
_get_cross_references,
_get_file_id,
_group_by_cluster_hint,
_split_group,
_topological_sort_with_levels,
_try_merge_singletons,
assign_clusters,
shard_fact_registry,
)
# ---------------------------------------------------------------------------
# Fixtures: sample data
# ---------------------------------------------------------------------------
def _make_entry(file_id: str, cluster_hint: str = "misc", cross_refs: list[str] | None = None) -> dict:
"""Helper to make a minimal manifest entry."""
entry = {"file_id": file_id, "cluster_hint": cluster_hint}
if cross_refs:
entry["cross_references"] = cross_refs
return entry
def _sample_fact_registry() -> dict:
"""A realistic fact registry for testing sharding."""
return {
"scenario_id": "dp_001",
"people": [
{"id": "john_doe", "full_name": "John Doe", "role": "CEO"},
{"id": "jane_smith", "full_name": "Jane Smith", "role": "CTO"},
],
"organizations": [
{"id": "acme_corp", "name": "Acme Corp", "type": "company"},
],
"dates": [
{"id": "date_kickoff", "date": "2024-01-15", "event": "Project kickoff", "files": ["f001", "f002"]},
{"id": "date_launch", "date": "2024-06-01", "event": "Product launch", "files": ["f005", "f006"]},
],
"financial": [
{"id": "budget_q1", "value": "$50,000.00", "description": "Q1 budget", "files": ["f001", "f003"]},
{"id": "budget_q2", "value": "$75,000.00", "description": "Q2 budget", "files": ["f004"]},
],
"references": [
{"id": "ref_contract", "value": "CTR-2024-001", "type": "contract", "files": ["f001", "f002"]},
],
"locations": [
{"id": "hq_office", "name": "HQ", "address": "123 Main St", "files": ["f001"]},
{"id": "branch_office", "name": "Branch", "address": "456 Oak Ave", "files": ["f005"]},
],
"domain_facts": [
{"id": "tech_stack", "category": "technical", "fact": "Uses Python 3.12", "files": ["f003", "f004"]},
],
"cross_references": [
{"source_file": "f001", "target_file": "f002", "fact_ids": ["date_kickoff"], "description": "kickoff ref"},
{"source_file": "f005", "target_file": "f006", "fact_ids": ["date_launch"], "description": "launch ref"},
],
}
# ---------------------------------------------------------------------------
# Tests: shard_fact_registry
# ---------------------------------------------------------------------------
class TestShardFactRegistry:
def test_global_categories_always_included(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f001"])
assert shard["people"] == registry["people"]
assert shard["organizations"] == registry["organizations"]
def test_scalar_fields_copied(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f001"])
assert shard["scenario_id"] == "dp_001"
def test_scoped_dates_filtered(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f001"])
# f001 is in date_kickoff but not date_launch
assert len(shard["dates"]) == 1
assert shard["dates"][0]["id"] == "date_kickoff"
def test_scoped_financial_filtered(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f004"])
assert len(shard["financial"]) == 1
assert shard["financial"][0]["id"] == "budget_q2"
def test_scoped_locations_filtered(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f005"])
assert len(shard["locations"]) == 1
assert shard["locations"][0]["id"] == "branch_office"
def test_cross_references_filtered(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f001"])
assert len(shard["cross_references"]) == 1
assert shard["cross_references"][0]["source_file"] == "f001"
def test_multiple_file_ids(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f001", "f005", "f006"])
# dates: date_kickoff (f001,f002) and date_launch (f005,f006)
assert len(shard["dates"]) == 2
# cross_references: both entries touch our files
assert len(shard["cross_references"]) == 2
def test_no_matching_files_returns_empty_scoped(self):
registry = _sample_fact_registry()
shard = shard_fact_registry(registry, ["f999"])
assert shard["dates"] == []
assert shard["financial"] == []
assert shard["cross_references"] == []
# Globals still present
assert len(shard["people"]) == 2
def test_empty_registry(self):
shard = shard_fact_registry({}, ["f001"])
assert shard == {}
def test_missing_files_key_in_entry(self):
registry = {
"dates": [{"id": "d1", "date": "2024-01-01"}], # no 'files' key
}
shard = shard_fact_registry(registry, ["f001"])
assert shard["dates"] == []
# ---------------------------------------------------------------------------
# Tests: helper functions
# ---------------------------------------------------------------------------
class TestHelpers:
def test_get_file_id(self):
assert _get_file_id({"file_id": "f001"}) == "f001"
assert _get_file_id({}) == ""
def test_get_cluster_hint_present(self):
assert _get_cluster_hint({"cluster_hint": "engineering"}) == "engineering"
def test_get_cluster_hint_missing(self):
assert _get_cluster_hint({}) == "misc"
def test_get_cluster_hint_empty_string(self):
assert _get_cluster_hint({"cluster_hint": ""}) == "misc"
def test_get_cluster_hint_none(self):
assert _get_cluster_hint({"cluster_hint": None}) == "misc"
def test_get_cross_references(self):
assert _get_cross_references({"cross_references": ["f002", "f003"]}) == ["f002", "f003"]
assert _get_cross_references({}) == []
assert _get_cross_references({"cross_references": "not_a_list"}) == []
class TestGroupByClusterHint:
def test_basic_grouping(self):
manifest = [
_make_entry("f001", "eng"),
_make_entry("f002", "eng"),
_make_entry("f003", "sales"),
]
groups = _group_by_cluster_hint(manifest)
assert len(groups) == 2
assert len(groups["eng"]) == 2
assert len(groups["sales"]) == 1
def test_missing_hint_goes_to_misc(self):
manifest = [{"file_id": "f001"}]
groups = _group_by_cluster_hint(manifest)
assert "misc" in groups
class TestBuildCrossRefGraph:
def test_basic_bidirectional(self):
manifest = [
_make_entry("f001", cross_refs=["f002"]),
_make_entry("f002"),
]
graph = _build_cross_ref_graph(manifest, {"f001", "f002"})
assert "f002" in graph["f001"]
assert "f001" in graph["f002"]
def test_warns_on_missing_ref(self):
manifest = [
_make_entry("f001", cross_refs=["f999"]),
]
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
_build_cross_ref_graph(manifest, {"f001"})
assert len(w) == 1
assert "f999" in str(w[0].message)
def test_self_reference_ignored(self):
manifest = [_make_entry("f001", cross_refs=["f001"])]
graph = _build_cross_ref_graph(manifest, {"f001"})
assert graph.get("f001", set()) == set()
class TestFindConnectedComponents:
def test_single_component(self):
adj = {"a": {"b"}, "b": {"a", "c"}, "c": {"b"}}
components = _find_connected_components(["a", "b", "c"], adj)
assert len(components) == 1
assert set(components[0]) == {"a", "b", "c"}
def test_two_components(self):
adj = {"a": {"b"}, "b": {"a"}}
components = _find_connected_components(["a", "b", "c"], adj)
assert len(components) == 2
def test_no_edges(self):
components = _find_connected_components(["a", "b", "c"], {})
assert len(components) == 3
# ---------------------------------------------------------------------------
# Tests: splitting and merging
# ---------------------------------------------------------------------------
class TestSplitGroup:
def test_no_split_needed(self):
entries = [_make_entry(f"f{i:03d}") for i in range(5)]
result = _split_group(entries, 8, {})
assert len(result) == 1
assert len(result[0]) == 5
def test_splits_oversize_group(self):
entries = [_make_entry(f"f{i:03d}") for i in range(12)]
result = _split_group(entries, 5, {})
for chunk in result:
assert len(chunk) <= 5
# All entries accounted for
all_ids = {_get_file_id(e) for chunk in result for e in chunk}
assert len(all_ids) == 12
def test_keeps_cross_refs_together(self):
entries = [_make_entry(f"f{i:03d}") for i in range(10)]
# f000 and f001 cross-reference each other
adjacency = {"f000": {"f001"}, "f001": {"f000"}}
result = _split_group(entries, 5, adjacency)
# f000 and f001 should be in the same chunk
for chunk in result:
ids = {_get_file_id(e) for e in chunk}
if "f000" in ids:
assert "f001" in ids
break
class TestTryMergeSingletons:
def test_singleton_merged_into_referenced_group(self):
groups = {
"eng": [_make_entry("f001"), _make_entry("f002")],
"lone": [_make_entry("f003", cross_refs=["f001"])],
}
adjacency = {"f003": {"f001"}, "f001": {"f003"}}
result = _try_merge_singletons(groups, 8, adjacency)
assert "lone" not in result
assert len(result["eng"]) == 3
def test_singleton_not_merged_if_target_full(self):
groups = {
"eng": [_make_entry(f"f{i:03d}") for i in range(8)], # already at max
"lone": [_make_entry("f100", cross_refs=["f001"])],
}
adjacency = {"f100": {"f001"}, "f001": {"f100"}}
result = _try_merge_singletons(groups, 8, adjacency)
assert "lone" in result
def test_singleton_without_refs_stays(self):
groups = {
"eng": [_make_entry("f001")],
"lone": [_make_entry("f002")],
}
result = _try_merge_singletons(groups, 8, {})
assert "lone" in result
# ---------------------------------------------------------------------------
# Tests: dependency graph and cycle detection
# ---------------------------------------------------------------------------
class TestBuildClusterDependencyGraph:
def test_cross_cluster_dependency(self):
manifest = [
_make_entry("f001", "eng", cross_refs=["f003"]),
_make_entry("f002", "eng"),
_make_entry("f003", "sales"),
]
clusters = {
"eng": [manifest[0], manifest[1]],
"sales": [manifest[2]],
}
file_to_cluster = {"f001": "eng", "f002": "eng", "f003": "sales"}
deps = _build_cluster_dependency_graph(
clusters, file_to_cluster, manifest, {"f001", "f002", "f003"}
)
assert "sales" in deps.get("eng", set())
def test_same_cluster_no_dependency(self):
manifest = [
_make_entry("f001", "eng", cross_refs=["f002"]),
_make_entry("f002", "eng"),
]
clusters = {"eng": manifest}
file_to_cluster = {"f001": "eng", "f002": "eng"}
deps = _build_cluster_dependency_graph(
clusters, file_to_cluster, manifest, {"f001", "f002"}
)
assert deps.get("eng", set()) == set()
class TestDetectAndMergeCycles:
def test_no_cycles(self):
groups = {"a": [_make_entry("f001")], "b": [_make_entry("f002")]}
deps = {"b": {"a"}}
new_groups, new_deps = _detect_and_merge_cycles(groups, deps)
assert set(new_groups.keys()) == {"a", "b"}
def test_two_node_cycle_merged(self):
groups = {
"a": [_make_entry("f001")],
"b": [_make_entry("f002")],
}
deps = {"a": {"b"}, "b": {"a"}}
new_groups, new_deps = _detect_and_merge_cycles(groups, deps)
# Should merge into one cluster
assert len(new_groups) == 1
merged_key = list(new_groups.keys())[0]
assert len(new_groups[merged_key]) == 2
def test_three_node_cycle_merged(self):
groups = {
"a": [_make_entry("f001")],
"b": [_make_entry("f002")],
"c": [_make_entry("f003")],
}
deps = {"a": {"b"}, "b": {"c"}, "c": {"a"}}
new_groups, new_deps = _detect_and_merge_cycles(groups, deps)
assert len(new_groups) == 1
def test_partial_cycle_with_external_dep(self):
groups = {
"a": [_make_entry("f001")],
"b": [_make_entry("f002")],
"c": [_make_entry("f003")],
}
# a <-> b form a cycle, c depends on a
deps = {"a": {"b"}, "b": {"a"}, "c": {"a"}}
new_groups, new_deps = _detect_and_merge_cycles(groups, deps)
assert len(new_groups) == 2
# c should depend on the merged a/b cluster
merged = [k for k in new_groups if k != "c"][0]
assert merged in new_deps.get("c", set())
# ---------------------------------------------------------------------------
# Tests: topological sort
# ---------------------------------------------------------------------------
class TestTopologicalSortWithLevels:
def test_no_deps(self):
result = _topological_sort_with_levels(["a", "b", "c"], {})
assert all(level == 0 for _, level in result)
assert len(result) == 3
def test_linear_chain(self):
# c depends on b, b depends on a
deps = {"c": {"b"}, "b": {"a"}}
result = _topological_sort_with_levels(["a", "b", "c"], deps)
levels = {cid: lvl for cid, lvl in result}
assert levels["a"] == 0
assert levels["b"] == 1
assert levels["c"] == 2
def test_diamond_shape(self):
# d depends on b and c, b and c depend on a
deps = {"b": {"a"}, "c": {"a"}, "d": {"b", "c"}}
result = _topological_sort_with_levels(["a", "b", "c", "d"], deps)
levels = {cid: lvl for cid, lvl in result}
assert levels["a"] == 0
assert levels["b"] == 1
assert levels["c"] == 1
assert levels["d"] == 2
def test_sorted_by_level_then_id(self):
deps = {"b": {"a"}, "c": {"a"}}
result = _topological_sort_with_levels(["c", "b", "a"], deps)
assert result[0] == ("a", 0)
# b and c are both level 1, sorted alphabetically
assert result[1] == ("b", 1)
assert result[2] == ("c", 1)
# ---------------------------------------------------------------------------
# Tests: assign_clusters (integration)
# ---------------------------------------------------------------------------
class TestAssignClusters:
def test_empty_manifest(self):
assert assign_clusters([], {}) == []
def test_basic_clustering(self):
manifest = [
_make_entry("f001", "eng"),
_make_entry("f002", "eng"),
_make_entry("f003", "sales"),
]
clusters = assign_clusters(manifest, {})
assert len(clusters) == 2
cluster_ids = {c.cluster_id for c in clusters}
assert "eng" in cluster_ids
assert "sales" in cluster_ids
def test_cluster_max_size_enforced(self):
manifest = [_make_entry(f"f{i:03d}", "big") for i in range(15)]
clusters = assign_clusters(manifest, {}, max_cluster_size=5)
for c in clusters:
assert len(c.file_entries) <= 5
total_files = sum(len(c.file_entries) for c in clusters)
assert total_files == 15
def test_cross_cluster_dependencies(self):
# Each cluster must have >1 file to avoid singleton merge
manifest = [
_make_entry("f001", "eng", cross_refs=["f003"]),
_make_entry("f002", "eng"),
_make_entry("f003", "sales"),
_make_entry("f004", "sales"),
]
clusters = assign_clusters(manifest, {})
eng_cluster = next(c for c in clusters if c.cluster_id == "eng")
# eng depends on sales because f001 references f003
assert "sales" in eng_cluster.depends_on
def test_levels_assigned_correctly(self):
# Each cluster needs >1 file to avoid singleton merge
manifest = [
_make_entry("f001", "base"),
_make_entry("f001b", "base"),
_make_entry("f002", "mid", cross_refs=["f001"]),
_make_entry("f002b", "mid"),
_make_entry("f003", "top", cross_refs=["f002"]),
_make_entry("f003b", "top"),
]
clusters = assign_clusters(manifest, {})
level_map = {c.cluster_id: c.level for c in clusters}
assert level_map["base"] == 0
assert level_map["mid"] == 1
assert level_map["top"] == 2
def test_clusters_ordered_by_level(self):
# Each cluster needs >1 file to avoid singleton merge
manifest = [
_make_entry("f001", "base"),
_make_entry("f001b", "base"),
_make_entry("f002", "mid", cross_refs=["f001"]),
_make_entry("f002b", "mid"),
_make_entry("f003", "top", cross_refs=["f002"]),
_make_entry("f003b", "top"),
]
clusters = assign_clusters(manifest, {})
levels = [c.level for c in clusters]
assert levels == sorted(levels)
def test_fact_sharding_integrated(self):
manifest = [
_make_entry("f001", "eng"),
_make_entry("f005", "ops"),
]
registry = _sample_fact_registry()
clusters = assign_clusters(manifest, registry)
eng = next(c for c in clusters if c.cluster_id == "eng")
ops = next(c for c in clusters if c.cluster_id == "ops")
# eng cluster (f001): should have date_kickoff, budget_q1, ref_contract, hq_office
eng_date_ids = {d["id"] for d in eng.fact_shard.get("dates", [])}
assert "date_kickoff" in eng_date_ids
assert "date_launch" not in eng_date_ids
# ops cluster (f005): should have date_launch, branch_office
ops_location_ids = {l["id"] for l in ops.fact_shard.get("locations", [])}
assert "branch_office" in ops_location_ids
assert "hq_office" not in ops_location_ids
def test_missing_cluster_hint_goes_to_misc(self):
manifest = [{"file_id": "f001"}, {"file_id": "f002"}]
clusters = assign_clusters(manifest, {})
assert len(clusters) == 1
assert clusters[0].cluster_id == "misc"
def test_circular_dependencies_merged(self):
manifest = [
_make_entry("f001", "alpha", cross_refs=["f003"]),
_make_entry("f002", "alpha"),
_make_entry("f003", "beta", cross_refs=["f001"]),
_make_entry("f004", "beta"),
]
clusters = assign_clusters(manifest, {})
# alpha and beta form a cycle — should be merged
assert len(clusters) == 1
assert len(clusters[0].file_entries) == 4
assert clusters[0].depends_on == []
def test_cross_ref_to_unknown_file_warns(self):
manifest = [
_make_entry("f001", "eng", cross_refs=["f999"]),
]
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
clusters = assign_clusters(manifest, {})
assert len(clusters) == 1
assert any("f999" in str(warning.message) for warning in w)
def test_singleton_merged_into_cross_ref_group(self):
manifest = [
_make_entry("f001", "eng"),
_make_entry("f002", "eng"),
_make_entry("f003", "lone", cross_refs=["f001"]),
]
clusters = assign_clusters(manifest, {}, max_cluster_size=8)
# f003 singleton should merge into eng
cluster_ids = {c.cluster_id for c in clusters}
assert "lone" not in cluster_ids
eng = next(c for c in clusters if c.cluster_id == "eng")
eng_file_ids = {_get_file_id(e) for e in eng.file_entries}
assert "f003" in eng_file_ids
def test_returns_cluster_dataclass(self):
manifest = [_make_entry("f001", "eng")]
clusters = assign_clusters(manifest, {})
assert len(clusters) == 1
c = clusters[0]
assert isinstance(c, Cluster)
assert c.cluster_id == "eng"
assert len(c.file_entries) == 1
assert isinstance(c.fact_shard, dict)
assert isinstance(c.depends_on, list)
assert isinstance(c.level, int)
def test_large_corpus_sharding(self):
"""Simulate a large corpus with 250+ files across many departments."""
manifest = []
departments = ["eng", "sales", "hr", "legal", "ops", "finance", "marketing", "support"]
fid = 1
for dept in departments:
for _ in range(32): # 32 files per dept = 256 total
manifest.append(_make_entry(f"f{fid:04d}", dept))
fid += 1
# Add some cross-refs
manifest[0]["cross_references"] = ["f0033"] # eng -> sales
manifest[64]["cross_references"] = ["f0001"] # hr -> eng
registry = {
"scenario_id": "dp_large",
"people": [{"id": "p1", "full_name": "Test Person"}],
"organizations": [],
"dates": [
{"id": f"d{i}", "date": "2024-01-01", "files": [f"f{i:04d}"]}
for i in range(1, 257)
],
}
clusters = assign_clusters(manifest, registry, max_cluster_size=8)
# Verify all files accounted for
total = sum(len(c.file_entries) for c in clusters)
assert total == 256
# Verify no cluster exceeds max size
for c in clusters:
assert len(c.file_entries) <= 8
# Verify fact shards are smaller than full registry
for c in clusters:
shard_date_count = len(c.fact_shard.get("dates", []))
assert shard_date_count <= len(c.file_entries)
# Verify levels are sorted
levels = [c.level for c in clusters]
assert levels == sorted(levels)

View file

@ -0,0 +1,634 @@
"""Tests for planner.py — the planning module for the eval corpus data generator."""
from __future__ import annotations
import asyncio
import json
from pathlib import Path
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
# Import the module under test
import planner
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
SAMPLE_SCENARIO_BLOCK = """\
dp_001: A small tech startup "Acme Labs" with 5 employees planning a product launch.
The CEO (Alice Smith), CTO (Bob Jones), and designer (Carol Lee) exchange emails,
Slack messages, and meeting notes over 3 weeks. Key decisions include pricing ($49/mo),
launch date (2025-03-15), and venue (Hilton Downtown, Room 301).
"""
SAMPLE_BRIEF = """\
# SCENARIO.md — dp_001
## 1. Overview
Scenario ID: dp_001
File count: 10
Time span: 2025-02-20 to 2025-03-15
Setting: Acme Labs, San Francisco
## 2. Cast of Characters
- alice_smith: Alice Smith, CEO, alice@acmelabs.com
- bob_jones: Bob Jones, CTO, bob@acmelabs.com
- carol_lee: Carol Lee, Designer, carol@acmelabs.com
## 3. Organizations
- Acme Labs: Tech startup, San Francisco
## 4. Timeline
- 2025-02-20: Kickoff meeting
- 2025-03-01: Pricing decision ($49/mo)
- 2025-03-15: Launch event at Hilton Downtown
## 5. Locked Facts Registry
- price_monthly: $49.00/mo
- launch_date: 2025-03-15
- venue: Hilton Downtown, Room 301
"""
SAMPLE_FACT_REGISTRY = {
"scenario_id": "dp_001",
"people": [
{
"id": "alice_smith",
"full_name": "Alice Smith",
"role": "CEO, Acme Labs",
"email": "alice@acmelabs.com",
"timezone": "America/Los_Angeles",
"location": "San Francisco, CA",
"traits": ["decisive", "formal writer"],
"writing_style": "Concise, professional emails",
"relationships": {"bob_jones": "direct report"},
},
{
"id": "bob_jones",
"full_name": "Bob Jones",
"role": "CTO, Acme Labs",
"email": "bob@acmelabs.com",
"timezone": "America/Los_Angeles",
"location": "San Francisco, CA",
"traits": ["technical", "uses jargon"],
"writing_style": "Detailed technical prose",
"relationships": {"alice_smith": "reports to"},
},
],
"organizations": [
{
"id": "acme_labs",
"name": "Acme Labs",
"type": "company",
"location": "San Francisco, CA",
"details": {"industry": "tech"},
}
],
"dates": [
{
"id": "kickoff_meeting",
"date": "2025-02-20",
"event": "Kickoff meeting",
"participants": ["alice_smith", "bob_jones"],
"files": ["f001", "f002"],
}
],
"financial": [
{
"id": "price_monthly",
"value": "$49.00",
"description": "Monthly subscription price",
"files": ["f003", "f005"],
}
],
"references": [],
"locations": [
{
"id": "hilton_downtown",
"name": "Hilton Downtown",
"address": "123 Market St, San Francisco, CA",
"type": "hotel",
"details": {"room": "301"},
"files": ["f004"],
}
],
"domain_facts": [],
"cross_references": [
{
"source_file": "f001",
"target_file": "f002",
"fact_ids": ["kickoff_meeting"],
"description": "Meeting notes reference email thread",
}
],
}
SAMPLE_MANIFEST_ENTRY = {
"file_id": "f001",
"path": "data/emails/kickoff_thread.eml",
"format": "email_thread",
"authors": ["alice_smith", "bob_jones"],
"date": "2025-02-20",
"target_tokens": [6000, 8000],
"locked_facts": ["kickoff_meeting", "price_monthly"],
"cross_references": ["f002", "f003"],
"cluster_hint": "communications",
"brief": "Email thread between Alice and Bob discussing the kickoff meeting. Contains pricing decisions.",
"tone": "formal",
"format_notes": "Standard email headers, threaded replies",
}
def _make_manifest(count: int) -> list[dict]:
"""Create a list of sample manifest entries."""
entries = []
for i in range(count):
entry = dict(SAMPLE_MANIFEST_ENTRY)
entry["file_id"] = f"f{i + 1:03d}"
entry["path"] = f"data/files/file_{i + 1:03d}.md"
entry["cross_references"] = []
entries.append(entry)
return entries
# ---------------------------------------------------------------------------
# Validation helpers tests
# ---------------------------------------------------------------------------
class TestValidateFactRegistry:
def test_valid_registry(self):
result = planner._validate_fact_registry(SAMPLE_FACT_REGISTRY)
assert result is SAMPLE_FACT_REGISTRY
def test_missing_people_key(self):
bad = {"organizations": [], "dates": []}
with pytest.raises(ValueError, match="missing required keys"):
planner._validate_fact_registry(bad)
def test_missing_organizations_key(self):
bad = {"people": [], "dates": []}
with pytest.raises(ValueError, match="missing required keys"):
planner._validate_fact_registry(bad)
def test_missing_dates_key(self):
bad = {"people": [], "organizations": []}
with pytest.raises(ValueError, match="missing required keys"):
planner._validate_fact_registry(bad)
def test_person_missing_id(self):
bad = {
"people": [{"full_name": "No ID"}],
"organizations": [],
"dates": [],
}
with pytest.raises(ValueError, match="missing 'id'"):
planner._validate_fact_registry(bad)
class TestValidateManifestEntry:
def test_valid_entry(self):
warnings = planner._validate_manifest_entry(SAMPLE_MANIFEST_ENTRY, 0)
assert warnings == []
def test_missing_fields(self):
warnings = planner._validate_manifest_entry({"file_id": "f001"}, 0)
assert any("missing fields" in w for w in warnings)
def test_tokens_out_of_range(self):
entry = dict(SAMPLE_MANIFEST_ENTRY)
entry["target_tokens"] = [1000, 20000]
warnings = planner._validate_manifest_entry(entry, 0)
assert any("outside [5000, 10000]" in w for w in warnings)
def test_tokens_min_gt_max(self):
entry = dict(SAMPLE_MANIFEST_ENTRY)
entry["target_tokens"] = [9000, 6000]
warnings = planner._validate_manifest_entry(entry, 0)
assert any("min > max" in w for w in warnings)
def test_tokens_valid_boundary(self):
entry = dict(SAMPLE_MANIFEST_ENTRY)
entry["target_tokens"] = [5000, 10000]
warnings = planner._validate_manifest_entry(entry, 0)
assert warnings == []
def test_tokens_malformed(self):
entry = dict(SAMPLE_MANIFEST_ENTRY)
entry["target_tokens"] = "not a list"
warnings = planner._validate_manifest_entry(entry, 0)
assert any("malformed" in w for w in warnings)
class TestRenumberManifest:
def test_sequential_renumbering(self):
entries = [
{"file_id": "f010", "cross_references": ["f020"]},
{"file_id": "f020", "cross_references": ["f010"]},
{"file_id": "f030", "cross_references": []},
]
result = planner._renumber_manifest(entries)
assert result[0]["file_id"] == "f001"
assert result[1]["file_id"] == "f002"
assert result[2]["file_id"] == "f003"
def test_cross_references_updated(self):
entries = [
{"file_id": "f010", "cross_references": ["f020", "f030"]},
{"file_id": "f020", "cross_references": ["f010"]},
{"file_id": "f030", "cross_references": ["f010"]},
]
result = planner._renumber_manifest(entries)
assert result[0]["cross_references"] == ["f002", "f003"]
assert result[1]["cross_references"] == ["f001"]
assert result[2]["cross_references"] == ["f001"]
def test_unknown_references_preserved(self):
entries = [
{"file_id": "f001", "cross_references": ["f999"]},
]
result = planner._renumber_manifest(entries)
# f999 is not in the manifest, so it stays as-is
assert result[0]["cross_references"] == ["f999"]
def test_empty_manifest(self):
result = planner._renumber_manifest([])
assert result == []
class TestValidateManifest:
def test_valid_manifest(self):
manifest = _make_manifest(3)
result = planner._validate_manifest(manifest)
assert len(result) == 3
def test_warnings_logged(self, caplog):
bad_entry = {"file_id": "f001"} # missing most fields
with caplog.at_level("WARNING"):
planner._validate_manifest([bad_entry])
assert "missing fields" in caplog.text
# ---------------------------------------------------------------------------
# Phase 1 tests
# ---------------------------------------------------------------------------
class TestGenerateScenarioBrief:
@pytest.mark.asyncio
async def test_generates_and_writes_file(self, tmp_path):
with patch("planner.llm_call", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = SAMPLE_BRIEF
result = await planner.generate_scenario_brief(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
assert result == SAMPLE_BRIEF
assert (tmp_path / "SCENARIO.md").exists()
assert (tmp_path / "SCENARIO.md").read_text() == SAMPLE_BRIEF
mock_llm.assert_called_once()
@pytest.mark.asyncio
async def test_resume_skips_existing(self, tmp_path):
# Pre-create SCENARIO.md
(tmp_path / "SCENARIO.md").write_text("existing brief")
with patch("planner.llm_call", new_callable=AsyncMock) as mock_llm:
result = await planner.generate_scenario_brief(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
assert result == "existing brief"
mock_llm.assert_not_called()
@pytest.mark.asyncio
async def test_passes_correct_model(self, tmp_path):
with patch("planner.llm_call", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = "brief"
await planner.generate_scenario_brief(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path, model="custom/model"
)
call_kwargs = mock_llm.call_args
assert call_kwargs.kwargs["model"] == "custom/model"
@pytest.mark.asyncio
async def test_uses_scenario_brief_system_prompt(self, tmp_path):
with patch("planner.llm_call", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = "brief"
await planner.generate_scenario_brief(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
call_kwargs = mock_llm.call_args
assert call_kwargs.kwargs["system"] == planner.SCENARIO_BRIEF_SYSTEM
# ---------------------------------------------------------------------------
# Phase 2 tests
# ---------------------------------------------------------------------------
class TestExtractFactRegistry:
@pytest.mark.asyncio
async def test_extracts_and_writes_file(self, tmp_path):
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = SAMPLE_FACT_REGISTRY
result = await planner.extract_fact_registry(SAMPLE_BRIEF, tmp_path)
assert result == SAMPLE_FACT_REGISTRY
assert (tmp_path / "facts.json").exists()
saved = json.loads((tmp_path / "facts.json").read_text())
assert saved["scenario_id"] == "dp_001"
mock_llm.assert_called_once()
@pytest.mark.asyncio
async def test_resume_skips_existing(self, tmp_path):
(tmp_path / "facts.json").write_text(json.dumps(SAMPLE_FACT_REGISTRY))
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
result = await planner.extract_fact_registry(SAMPLE_BRIEF, tmp_path)
assert result["scenario_id"] == "dp_001"
mock_llm.assert_not_called()
@pytest.mark.asyncio
async def test_unwraps_nested_response(self, tmp_path):
"""LLM might return {"result": {actual registry}}."""
wrapped = {"result": SAMPLE_FACT_REGISTRY}
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = wrapped
result = await planner.extract_fact_registry(SAMPLE_BRIEF, tmp_path)
assert result["scenario_id"] == "dp_001"
@pytest.mark.asyncio
async def test_validation_fails_on_bad_registry(self, tmp_path):
bad_registry = {"not_valid": True}
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = bad_registry
with pytest.raises(ValueError, match="missing required keys"):
await planner.extract_fact_registry(SAMPLE_BRIEF, tmp_path)
@pytest.mark.asyncio
async def test_uses_fact_registry_system_prompt(self, tmp_path):
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = SAMPLE_FACT_REGISTRY
await planner.extract_fact_registry(SAMPLE_BRIEF, tmp_path)
call_kwargs = mock_llm.call_args
assert call_kwargs.kwargs["system"] == planner.FACT_REGISTRY_SYSTEM
# ---------------------------------------------------------------------------
# Phase 3 tests
# ---------------------------------------------------------------------------
class TestGenerateManifest:
@pytest.mark.asyncio
async def test_small_manifest_single_call(self, tmp_path):
manifest = _make_manifest(10)
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = manifest
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 10, tmp_path
)
assert len(result) == 10
assert result[0]["file_id"] == "f001"
assert (tmp_path / "manifest.json").exists()
mock_llm.assert_called_once()
@pytest.mark.asyncio
async def test_small_manifest_unwraps_dict(self, tmp_path):
"""LLM might return {"files": [...]} instead of bare array."""
manifest = _make_manifest(5)
wrapped = {"files": manifest}
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.return_value = wrapped
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 5, tmp_path
)
assert len(result) == 5
@pytest.mark.asyncio
async def test_resume_skips_existing(self, tmp_path):
manifest = _make_manifest(10)
(tmp_path / "manifest.json").write_text(json.dumps(manifest))
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 10, tmp_path
)
assert len(result) == 10
mock_llm.assert_not_called()
@pytest.mark.asyncio
async def test_resume_unwraps_files_key(self, tmp_path):
manifest = _make_manifest(5)
(tmp_path / "manifest.json").write_text(json.dumps({"files": manifest}))
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 5, tmp_path
)
assert len(result) == 5
mock_llm.assert_not_called()
@pytest.mark.asyncio
async def test_large_manifest_chunked(self, tmp_path):
"""Corpora > 50 files should use chunked generation."""
outline = {
"sections": [
{
"name": "Section A",
"cluster_hint": "section_a",
"file_count": 30,
"description": "First section",
},
{
"name": "Section B",
"cluster_hint": "section_b",
"file_count": 30,
"description": "Second section",
},
]
}
section_a_entries = _make_manifest(30)
section_b_entries = _make_manifest(30)
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
# First call returns outline, next two return section entries
mock_llm.side_effect = [outline, section_a_entries, section_b_entries]
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 60, tmp_path
)
assert len(result) == 60
# Should have been called 3 times: outline + 2 sections
assert mock_llm.call_count == 3
# Verify sequential numbering
assert result[0]["file_id"] == "f001"
assert result[29]["file_id"] == "f030"
assert result[30]["file_id"] == "f031"
assert result[59]["file_id"] == "f060"
@pytest.mark.asyncio
async def test_large_manifest_adjusts_file_count(self, tmp_path):
"""If outline section counts don't add up, the last section is adjusted."""
outline = {
"sections": [
{
"name": "Section A",
"cluster_hint": "section_a",
"file_count": 25,
"description": "First section",
},
{
"name": "Section B",
"cluster_hint": "section_b",
"file_count": 24, # total=49, should be 55
"description": "Second section",
},
]
}
section_a_entries = _make_manifest(25)
section_b_entries = _make_manifest(30) # adjusted to 30
with patch("planner.llm_call_json", new_callable=AsyncMock) as mock_llm:
mock_llm.side_effect = [outline, section_a_entries, section_b_entries]
result = await planner.generate_manifest(
SAMPLE_BRIEF, SAMPLE_FACT_REGISTRY, 55, tmp_path
)
# Outline call should have adjusted Section B to 30
assert mock_llm.call_count == 3
# ---------------------------------------------------------------------------
# Orchestrator tests
# ---------------------------------------------------------------------------
class TestRunPlanning:
@pytest.mark.asyncio
async def test_runs_all_three_phases(self, tmp_path):
manifest = _make_manifest(10)
with patch("planner.llm_call", new_callable=AsyncMock) as mock_text, \
patch("planner.llm_call_json", new_callable=AsyncMock) as mock_json:
mock_text.return_value = SAMPLE_BRIEF
mock_json.side_effect = [SAMPLE_FACT_REGISTRY, manifest]
brief, facts, result_manifest = await planner.run_planning(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
assert brief == SAMPLE_BRIEF
assert facts == SAMPLE_FACT_REGISTRY
assert len(result_manifest) == 10
assert (tmp_path / "SCENARIO.md").exists()
assert (tmp_path / "facts.json").exists()
assert (tmp_path / "manifest.json").exists()
@pytest.mark.asyncio
async def test_creates_output_dir(self, tmp_path):
out = tmp_path / "nested" / "output"
manifest = _make_manifest(5)
with patch("planner.llm_call", new_callable=AsyncMock) as mock_text, \
patch("planner.llm_call_json", new_callable=AsyncMock) as mock_json:
mock_text.return_value = SAMPLE_BRIEF
mock_json.side_effect = [SAMPLE_FACT_REGISTRY, manifest]
await planner.run_planning(SAMPLE_SCENARIO_BLOCK, 5, out)
assert out.exists()
assert (out / "SCENARIO.md").exists()
@pytest.mark.asyncio
async def test_resumes_from_phase2(self, tmp_path):
"""If SCENARIO.md exists, skip Phase 1 and continue."""
(tmp_path / "SCENARIO.md").write_text(SAMPLE_BRIEF)
manifest = _make_manifest(10)
with patch("planner.llm_call", new_callable=AsyncMock) as mock_text, \
patch("planner.llm_call_json", new_callable=AsyncMock) as mock_json:
mock_json.side_effect = [SAMPLE_FACT_REGISTRY, manifest]
brief, facts, result_manifest = await planner.run_planning(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
# Phase 1 LLM call should NOT have been made
mock_text.assert_not_called()
assert brief == SAMPLE_BRIEF
@pytest.mark.asyncio
async def test_resumes_from_phase3(self, tmp_path):
"""If SCENARIO.md and facts.json exist, skip Phases 1 and 2."""
(tmp_path / "SCENARIO.md").write_text(SAMPLE_BRIEF)
(tmp_path / "facts.json").write_text(json.dumps(SAMPLE_FACT_REGISTRY))
manifest = _make_manifest(10)
with patch("planner.llm_call", new_callable=AsyncMock) as mock_text, \
patch("planner.llm_call_json", new_callable=AsyncMock) as mock_json:
mock_json.side_effect = [manifest]
brief, facts, result_manifest = await planner.run_planning(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
mock_text.assert_not_called()
# Only one json call (manifest), not two
assert mock_json.call_count == 1
@pytest.mark.asyncio
async def test_full_resume(self, tmp_path):
"""If all artifacts exist, no LLM calls made at all."""
(tmp_path / "SCENARIO.md").write_text(SAMPLE_BRIEF)
(tmp_path / "facts.json").write_text(json.dumps(SAMPLE_FACT_REGISTRY))
manifest = _make_manifest(10)
(tmp_path / "manifest.json").write_text(json.dumps(manifest))
with patch("planner.llm_call", new_callable=AsyncMock) as mock_text, \
patch("planner.llm_call_json", new_callable=AsyncMock) as mock_json:
brief, facts, result_manifest = await planner.run_planning(
SAMPLE_SCENARIO_BLOCK, 10, tmp_path
)
mock_text.assert_not_called()
mock_json.assert_not_called()
assert brief == SAMPLE_BRIEF
assert len(result_manifest) == 10
# ---------------------------------------------------------------------------
# Constants tests
# ---------------------------------------------------------------------------
class TestConstants:
def test_large_corpus_threshold(self):
assert planner.LARGE_CORPUS_THRESHOLD == 50
def test_chunk_size(self):
assert planner.CHUNK_SIZE == 30

View file

@ -0,0 +1,504 @@
"""Tests for questions.py — Phase 7: Eval Question Generation."""
from __future__ import annotations
import asyncio
import json
from pathlib import Path
from unittest.mock import AsyncMock, patch
import pytest
from questions import (
MAX_SCENARIO_TOKENS,
VALID_FAMILIES,
_build_fact_summary,
_build_manifest_summary,
_build_scenario_summary,
_sample_file_excerpts,
_truncate_to_tokens,
_validate_questions,
generate_questions,
)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def tmp_corpus(tmp_path: Path):
"""Create a minimal corpus directory with sample files."""
data_dir = tmp_path / "data" / "emails"
data_dir.mkdir(parents=True)
(data_dir / "test.md").write_text("This is a test email about project updates.")
contracts_dir = tmp_path / "data" / "contracts"
contracts_dir.mkdir(parents=True)
(contracts_dir / "contract.md").write_text("This is a legal contract between parties.")
return tmp_path
@pytest.fixture
def sample_manifest():
return [
{
"file_id": "f001",
"path": "data/emails/test.md",
"format": "email_thread",
"authors": ["person_alice"],
"date": "2026-04-22",
"target_tokens": [5000, 10000],
"locked_facts": ["fin_budget", "ref_booking"],
"cross_references": ["f002"],
"cluster_hint": "emails",
"brief": "Email thread discussing Q1 budget allocation and hotel booking",
"tone": "casual",
"format_notes": "",
},
{
"file_id": "f002",
"path": "data/contracts/contract.md",
"format": "legal_contract",
"authors": ["person_bob"],
"date": "2026-03-15",
"target_tokens": [5000, 10000],
"locked_facts": ["person_alice"],
"cross_references": ["f001"],
"cluster_hint": "legal",
"brief": "Legal contract for vendor services signed by Alice Johnson",
"tone": "formal",
"format_notes": "",
},
]
@pytest.fixture
def sample_fact_registry():
return {
"people": [
{
"id": "person_alice",
"full_name": "Alice Johnson",
"role": "CTO, Acme Corp",
},
{
"id": "person_bob",
"full_name": "Bob Williams",
"role": "General Counsel, Acme Corp",
},
],
"organizations": [
{"id": "org_acme", "name": "Acme Corp", "type": "company"},
],
"financial": [
{
"id": "fin_budget",
"value": "$2,034.50",
"description": "Q1 marketing budget",
"files": ["f001"],
},
],
"references": [
{
"id": "ref_booking",
"value": "BK-2026-0422",
"type": "booking",
"description": "Hotel booking reference",
"files": ["f001"],
},
],
"dates": [
{
"id": "date_deadline",
"date": "2026-04-22",
"event": "Project deadline",
"files": ["f001"],
},
],
"locations": [],
"domain_facts": [
{
"id": "domain_spec",
"category": "technical",
"fact": "The system uses PostgreSQL 15 with pgvector for embeddings",
"files": ["f001"],
},
],
"cross_references": [
{
"source_file": "f001",
"target_file": "f002",
"fact_ids": ["fin_budget"],
"description": "Contract references the budget",
},
],
}
@pytest.fixture
def sample_scenario_brief():
return (
"# Scenario: Acme Corp Q1 Planning\n\n"
"Acme Corp is a mid-size tech company planning their Q1 budget. "
"Alice Johnson (CTO) is coordinating with Bob Williams (General Counsel) "
"on vendor contracts and travel arrangements.\n\n"
"## Timeline\n"
"- 2026-03-15: Contract signing\n"
"- 2026-04-22: Project deadline\n"
)
# ---------------------------------------------------------------------------
# Tests: _truncate_to_tokens
# ---------------------------------------------------------------------------
class TestTruncateToTokens:
def test_short_text_unchanged(self):
text = "Short text."
result = _truncate_to_tokens(text, 1000)
assert result == text
def test_long_text_truncated(self):
text = "Word " * 5000 # ~5000 tokens
result = _truncate_to_tokens(text, 100)
assert len(result) < len(text)
assert result.endswith("[… truncated …]")
def test_truncation_at_boundary(self):
text = "First sentence. Second sentence. Third sentence. Fourth sentence."
result = _truncate_to_tokens(text, 2)
# Should be truncated
assert len(result) < len(text)
# ---------------------------------------------------------------------------
# Tests: _build_scenario_summary
# ---------------------------------------------------------------------------
class TestBuildScenarioSummary:
def test_short_brief_unchanged(self, sample_scenario_brief):
result = _build_scenario_summary(sample_scenario_brief)
# Short brief should not be truncated
assert "Acme Corp" in result
def test_long_brief_truncated(self):
long_brief = "Context. " * 10000
result = _build_scenario_summary(long_brief)
assert len(result) < len(long_brief)
# ---------------------------------------------------------------------------
# Tests: _build_fact_summary
# ---------------------------------------------------------------------------
class TestBuildFactSummary:
def test_includes_people(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "Alice Johnson" in result
assert "Bob Williams" in result
def test_includes_financial(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "$2,034.50" in result
def test_includes_references(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "BK-2026-0422" in result
def test_includes_dates(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "2026-04-22" in result
def test_includes_domain_facts(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "PostgreSQL" in result
def test_includes_cross_references_count(self, sample_fact_registry):
result = _build_fact_summary(sample_fact_registry)
assert "1 connections" in result
def test_empty_registry(self):
result = _build_fact_summary({})
assert result == ""
def test_many_domain_facts_truncated(self):
registry = {
"domain_facts": [
{"id": f"df_{i}", "fact": f"Fact number {i} about something"} for i in range(20)
],
}
result = _build_fact_summary(registry)
assert "and 10 more" in result
# ---------------------------------------------------------------------------
# Tests: _build_manifest_summary
# ---------------------------------------------------------------------------
class TestBuildManifestSummary:
def test_includes_file_ids(self, sample_manifest):
result = _build_manifest_summary(sample_manifest)
assert "f001" in result
assert "f002" in result
def test_includes_paths(self, sample_manifest):
result = _build_manifest_summary(sample_manifest)
assert "data/emails/test.md" in result
assert "data/contracts/contract.md" in result
def test_includes_formats(self, sample_manifest):
result = _build_manifest_summary(sample_manifest)
assert "email_thread" in result
assert "legal_contract" in result
def test_includes_briefs(self, sample_manifest):
result = _build_manifest_summary(sample_manifest)
assert "budget" in result.lower()
def test_includes_locked_facts(self, sample_manifest):
result = _build_manifest_summary(sample_manifest)
assert "fin_budget" in result
def test_empty_manifest(self):
result = _build_manifest_summary([])
assert result == ""
# ---------------------------------------------------------------------------
# Tests: _sample_file_excerpts
# ---------------------------------------------------------------------------
class TestSampleFileExcerpts:
def test_samples_existing_files(self, tmp_corpus, sample_manifest):
result = _sample_file_excerpts(tmp_corpus, sample_manifest)
assert "f001" in result
assert "test email" in result.lower()
def test_empty_manifest(self, tmp_corpus):
result = _sample_file_excerpts(tmp_corpus, [])
assert result == ""
def test_missing_files_skipped(self, tmp_corpus):
manifest = [
{"file_id": "f999", "path": "data/nonexistent/file.md"},
]
result = _sample_file_excerpts(tmp_corpus, manifest)
assert result == ""
# ---------------------------------------------------------------------------
# Tests: _validate_questions
# ---------------------------------------------------------------------------
class TestValidateQuestions:
def test_valid_questions_pass(self, sample_manifest):
questions = [
{
"id": "q01",
"family": "single_hop",
"prompt": "What is the Q1 budget?",
"gold_file_ids": ["data/emails/test.md"],
"gold_answer": "$2,034.50",
},
]
result = _validate_questions(questions, sample_manifest)
assert len(result) == 1
assert result[0]["family"] == "single_hop"
def test_missing_required_fields_skipped(self, sample_manifest):
questions = [
{"id": "q01", "family": "single_hop"}, # missing prompt, gold_file_ids, gold_answer
]
result = _validate_questions(questions, sample_manifest)
assert len(result) == 0
def test_file_id_normalized_to_path(self, sample_manifest):
questions = [
{
"id": "q01",
"family": "single_hop",
"prompt": "What is the Q1 budget?",
"gold_file_ids": ["f001"], # file_id instead of path
"gold_answer": "$2,034.50",
},
]
result = _validate_questions(questions, sample_manifest)
assert result[0]["gold_file_ids"] == ["data/emails/test.md"]
def test_family_normalization(self, sample_manifest):
questions = [
{
"id": "q01",
"family": "Multi-Hop",
"prompt": "What is the total cost?",
"gold_file_ids": ["data/emails/test.md"],
"gold_answer": "$5,000",
},
]
result = _validate_questions(questions, sample_manifest)
assert result[0]["family"] == "multi_hop"
def test_unknown_file_id_preserved(self, sample_manifest):
questions = [
{
"id": "q01",
"family": "single_hop",
"prompt": "Test?",
"gold_file_ids": ["data/unknown/file.md"],
"gold_answer": "answer",
},
]
result = _validate_questions(questions, sample_manifest)
assert result[0]["gold_file_ids"] == ["data/unknown/file.md"]
# ---------------------------------------------------------------------------
# Tests: generate_questions (integration)
# ---------------------------------------------------------------------------
class TestGenerateQuestions:
@patch("questions.llm_call_json", new_callable=AsyncMock)
def test_generates_and_writes(
self,
mock_llm,
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
):
"""Integration test: generates questions and writes question.json."""
mock_llm.return_value = [
{
"id": "q01",
"family": "single_hop",
"prompt": "What is the Q1 marketing budget?",
"gold_file_ids": ["data/emails/test.md"],
"gold_answer": "$2,034.50",
},
{
"id": "q02",
"family": "multi_hop",
"prompt": "Who signed the contract and what was the budget?",
"gold_file_ids": ["data/emails/test.md", "data/contracts/contract.md"],
"gold_answer": "Alice Johnson signed; budget was $2,034.50",
},
]
result = asyncio.run(
generate_questions(
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
)
)
assert len(result) == 2
assert result[0]["id"] == "q01"
assert result[1]["family"] == "multi_hop"
# Should write to disk
output_path = tmp_corpus / "question.json"
assert output_path.exists()
written = json.loads(output_path.read_text())
assert len(written) == 2
@patch("questions.llm_call_json", new_callable=AsyncMock)
def test_handles_wrapped_response(
self,
mock_llm,
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
):
"""LLM may return questions wrapped in a dict."""
mock_llm.return_value = {
"questions": [
{
"id": "q01",
"family": "single_hop",
"prompt": "Test question?",
"gold_file_ids": ["f001"],
"gold_answer": "Answer",
},
]
}
result = asyncio.run(
generate_questions(
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
)
)
assert len(result) == 1
# f001 should be normalized to path
assert result[0]["gold_file_ids"] == ["data/emails/test.md"]
def test_resume_support(
self,
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
):
"""If question.json already exists, it should be returned without LLM call."""
output_path = tmp_corpus / "question.json"
existing = [{"id": "q01", "family": "single_hop", "prompt": "Existing?", "gold_file_ids": [], "gold_answer": "yes"}]
output_path.write_text(json.dumps(existing))
result = asyncio.run(
generate_questions(
tmp_corpus,
sample_scenario_brief,
sample_fact_registry,
sample_manifest,
)
)
# Should return the existing content (as string since read_text returns string)
assert "Existing?" in str(result)
# ---------------------------------------------------------------------------
# Tests: prompts/questions.py
# ---------------------------------------------------------------------------
class TestQuestionPrompts:
def test_system_prompt_exists(self):
from prompts.questions import QUESTION_GEN_SYSTEM
assert "single_hop" in QUESTION_GEN_SYSTEM
assert "multi_hop" in QUESTION_GEN_SYSTEM
assert "format_spanning" in QUESTION_GEN_SYSTEM
assert "edit_then_recall" in QUESTION_GEN_SYSTEM
def test_prompt_template_has_placeholders(self):
from prompts.questions import QUESTION_GEN_PROMPT
assert "{scenario_summary}" in QUESTION_GEN_PROMPT
assert "{fact_summary}" in QUESTION_GEN_PROMPT
assert "{manifest_summary}" in QUESTION_GEN_PROMPT
def test_prompt_template_formats(self):
from prompts.questions import QUESTION_GEN_PROMPT
result = QUESTION_GEN_PROMPT.format(
scenario_summary="Test scenario",
fact_summary="Test facts",
manifest_summary="Test manifest",
)
assert "Test scenario" in result
assert "Test facts" in result
assert "gold_file_ids" in result

View file

@ -0,0 +1,509 @@
"""Tests for validator.py — Phase 6: Cross-Reference & Consistency Audit."""
from __future__ import annotations
import asyncio
import json
from pathlib import Path
from unittest.mock import AsyncMock, patch
import pytest
from validator import (
TOKEN_MAX,
TOKEN_MIN,
ValidationIssue,
ValidationReport,
_check_cross_references,
_check_file_existence,
_check_locked_facts,
_check_name_consistency,
_check_token_counts,
_normalize_date,
validate_corpus,
repair_files,
)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def tmp_corpus(tmp_path: Path):
"""Create a minimal corpus directory with a single valid file."""
data_dir = tmp_path / "data" / "emails"
data_dir.mkdir(parents=True)
# Write a file with known content (~100 tokens is enough for basic tests)
content = "Hello World. " * 500 # roughly 500 tokens — below minimum
(data_dir / "test.md").write_text(content)
return tmp_path
@pytest.fixture
def sample_manifest():
"""A minimal manifest for testing."""
return [
{
"file_id": "f001",
"path": "data/emails/test.md",
"format": "markdown_prose",
"authors": ["person_alice"],
"date": "2026-04-22",
"target_tokens": [5000, 10000],
"locked_facts": ["fin_budget", "ref_booking"],
"cross_references": ["f002"],
"cluster_hint": "emails",
"brief": "Test email file",
"tone": "casual",
"format_notes": "",
},
{
"file_id": "f002",
"path": "data/contracts/contract.md",
"format": "legal_contract",
"authors": ["person_bob"],
"date": "2026-03-15",
"target_tokens": [5000, 10000],
"locked_facts": ["person_alice"],
"cross_references": ["f001"],
"cluster_hint": "legal",
"brief": "Legal contract",
"tone": "formal",
"format_notes": "",
},
]
@pytest.fixture
def sample_fact_registry():
"""A minimal fact registry for testing."""
return {
"people": [
{
"id": "person_alice",
"full_name": "Alice Johnson",
"role": "CTO, Acme Corp",
"email": "alice@acme.com",
},
{
"id": "person_bob",
"full_name": "Bob Williams",
"role": "General Counsel, Acme Corp",
"email": "bob@acme.com",
},
],
"organizations": [
{
"id": "org_acme",
"name": "Acme Corp",
"type": "company",
"location": "San Francisco, CA",
},
],
"financial": [
{
"id": "fin_budget",
"value": "$2,034.50",
"description": "Q1 marketing budget",
"files": ["f001"],
},
],
"references": [
{
"id": "ref_booking",
"value": "BK-2026-0422",
"type": "booking",
"description": "Hotel booking reference",
"files": ["f001"],
},
],
"dates": [
{
"id": "date_deadline",
"date": "2026-04-22",
"event": "Project deadline",
"files": ["f001"],
},
],
"locations": [],
"domain_facts": [],
"cross_references": [],
}
# ---------------------------------------------------------------------------
# Tests: ValidationReport
# ---------------------------------------------------------------------------
class TestValidationReport:
def test_errors_and_warnings(self):
issues = [
ValidationIssue("f001", "token_count", "error", "too short"),
ValidationIssue("f002", "name_inconsistency", "warning", "name mismatch"),
ValidationIssue("f003", "missing_fact", "error", "fact missing"),
]
report = ValidationReport(total_files=3, files_checked=3, issues=issues)
assert len(report.errors) == 2
assert len(report.warnings) == 1
assert report.errors[0].file_id == "f001"
assert report.errors[1].file_id == "f003"
assert report.warnings[0].file_id == "f002"
def test_empty_report(self):
report = ValidationReport(total_files=5, files_checked=5)
assert report.errors == []
assert report.warnings == []
# ---------------------------------------------------------------------------
# Tests: _normalize_date
# ---------------------------------------------------------------------------
class TestNormalizeDate:
def test_iso_date(self):
variants = _normalize_date("2026-04-22")
assert "2026-04-22" in variants
assert "April 22, 2026" in variants
assert "Apr 22, 2026" in variants
assert "04/22/2026" in variants
assert "22 April 2026" in variants
def test_january(self):
variants = _normalize_date("2025-01-05")
assert "January 5, 2025" in variants
assert "Jan 5, 2025" in variants
assert "01/05/2025" in variants
def test_non_date_string(self):
variants = _normalize_date("not-a-date")
assert variants == ["not-a-date"]
def test_december(self):
variants = _normalize_date("2024-12-31")
assert "December 31, 2024" in variants
assert "Dec 31, 2024" in variants
# ---------------------------------------------------------------------------
# Tests: _check_file_existence
# ---------------------------------------------------------------------------
class TestCheckFileExistence:
def test_existing_file(self, tmp_corpus, sample_manifest):
# Only f001 exists
issues = _check_file_existence(tmp_corpus, [sample_manifest[0]])
assert len(issues) == 0
def test_missing_file(self, tmp_corpus, sample_manifest):
# f002 doesn't exist
issues = _check_file_existence(tmp_corpus, [sample_manifest[1]])
assert len(issues) == 1
assert issues[0].issue_type == "file_missing"
assert issues[0].severity == "error"
def test_mixed_existence(self, tmp_corpus, sample_manifest):
issues = _check_file_existence(tmp_corpus, sample_manifest)
# f001 exists, f002 does not
assert len(issues) == 1
assert issues[0].file_id == "f002"
# ---------------------------------------------------------------------------
# Tests: _check_token_counts
# ---------------------------------------------------------------------------
class TestCheckTokenCounts:
def test_file_below_minimum(self, tmp_corpus, sample_manifest):
# The test file has ~500 tokens, well below TOKEN_MIN
issues, token_map, stats = _check_token_counts(tmp_corpus, [sample_manifest[0]])
assert len(issues) == 1
assert issues[0].issue_type == "token_count"
assert issues[0].severity == "error"
assert "below minimum" in issues[0].description
def test_file_in_range(self, tmp_corpus, sample_manifest):
# Write a file with enough tokens
file_path = tmp_corpus / "data" / "emails" / "test.md"
content = "The quick brown fox jumps over the lazy dog. " * 1200 # ~12k tokens → adjust
file_path.write_text(content)
issues, token_map, stats = _check_token_counts(tmp_corpus, [sample_manifest[0]])
token_count = token_map.get("f001", 0)
if TOKEN_MIN <= token_count <= TOKEN_MAX:
assert len(issues) == 0
# If our estimate is wrong, just verify the check ran
assert "f001" in token_map
def test_stats_computed(self, tmp_corpus, sample_manifest):
issues, token_map, stats = _check_token_counts(tmp_corpus, [sample_manifest[0]])
assert "min" in stats
assert "max" in stats
assert "mean" in stats
assert "median" in stats
def test_missing_file_skipped(self, tmp_corpus, sample_manifest):
# f002 doesn't exist — should be silently skipped
issues, token_map, stats = _check_token_counts(tmp_corpus, [sample_manifest[1]])
assert len(issues) == 0
assert "f002" not in token_map
# ---------------------------------------------------------------------------
# Tests: _check_locked_facts
# ---------------------------------------------------------------------------
class TestCheckLockedFacts:
def test_all_facts_present(self, tmp_corpus, sample_manifest, sample_fact_registry):
# Write content that contains all locked facts for f001
file_path = tmp_corpus / "data" / "emails" / "test.md"
content = (
"The Q1 marketing budget is $2,034.50 and the booking reference "
"is BK-2026-0422. Alice Johnson confirmed on April 22, 2026."
)
file_path.write_text(content)
issues = _check_locked_facts(tmp_corpus, [sample_manifest[0]], sample_fact_registry)
# Should have no errors for f001 — all locked facts are present
f001_errors = [i for i in issues if i.file_id == "f001" and i.severity == "error"]
assert len(f001_errors) == 0
def test_missing_financial_fact(self, tmp_corpus, sample_manifest, sample_fact_registry):
file_path = tmp_corpus / "data" / "emails" / "test.md"
content = "The booking reference is BK-2026-0422. No budget info here."
file_path.write_text(content)
issues = _check_locked_facts(tmp_corpus, [sample_manifest[0]], sample_fact_registry)
error_ids = [i.details.get("fact_id") for i in issues if i.severity == "error"]
assert "fin_budget" in error_ids
def test_missing_reference_fact(self, tmp_corpus, sample_manifest, sample_fact_registry):
file_path = tmp_corpus / "data" / "emails" / "test.md"
content = "The budget is $2,034.50 but no booking reference here."
file_path.write_text(content)
issues = _check_locked_facts(tmp_corpus, [sample_manifest[0]], sample_fact_registry)
error_ids = [i.details.get("fact_id") for i in issues if i.severity == "error"]
assert "ref_booking" in error_ids
def test_date_variant_matching(self, tmp_corpus, sample_fact_registry):
"""Check that date facts match any common format variant."""
manifest_with_date = [
{
"file_id": "f001",
"path": "data/emails/test.md",
"locked_facts": ["date_deadline"],
"cross_references": [],
},
]
# Test ISO format
file_path = tmp_corpus / "data" / "emails" / "test.md"
file_path.write_text("Deadline is 2026-04-22.")
issues = _check_locked_facts(tmp_corpus, manifest_with_date, sample_fact_registry)
date_errors = [i for i in issues if i.details.get("fact_id") == "date_deadline" and i.severity == "error"]
assert len(date_errors) == 0
# Test natural format
file_path.write_text("Deadline is April 22, 2026.")
issues = _check_locked_facts(tmp_corpus, manifest_with_date, sample_fact_registry)
date_errors = [i for i in issues if i.details.get("fact_id") == "date_deadline" and i.severity == "error"]
assert len(date_errors) == 0
# Test abbreviated format
file_path.write_text("Deadline is Apr 22, 2026.")
issues = _check_locked_facts(tmp_corpus, manifest_with_date, sample_fact_registry)
date_errors = [i for i in issues if i.details.get("fact_id") == "date_deadline" and i.severity == "error"]
assert len(date_errors) == 0
def test_person_name_check(self, tmp_corpus, sample_manifest, sample_fact_registry):
"""Check that person names are found case-insensitively."""
# f002 has locked_facts: ["person_alice"]
contract_dir = tmp_corpus / "data" / "contracts"
contract_dir.mkdir(parents=True, exist_ok=True)
(contract_dir / "contract.md").write_text("Contract signed by alice johnson.")
issues = _check_locked_facts(tmp_corpus, [sample_manifest[1]], sample_fact_registry)
person_errors = [i for i in issues if i.details.get("fact_id") == "person_alice" and i.severity == "error"]
assert len(person_errors) == 0
def test_unknown_fact_id_warns(self, tmp_corpus, sample_fact_registry):
manifest_with_unknown = [
{
"file_id": "f001",
"path": "data/emails/test.md",
"locked_facts": ["nonexistent_fact"],
"cross_references": [],
},
]
issues = _check_locked_facts(tmp_corpus, manifest_with_unknown, sample_fact_registry)
warnings = [i for i in issues if i.severity == "warning"]
assert len(warnings) == 1
assert "not found in fact registry" in warnings[0].description
# ---------------------------------------------------------------------------
# Tests: _check_name_consistency
# ---------------------------------------------------------------------------
class TestCheckNameConsistency:
def test_consistent_names(self, tmp_corpus, sample_manifest, sample_fact_registry):
file_path = tmp_corpus / "data" / "emails" / "test.md"
file_path.write_text("Alice Johnson sent a message. Johnson approved the plan.")
issues = _check_name_consistency(tmp_corpus, [sample_manifest[0]], sample_fact_registry)
# "Johnson" appears and "Alice Johnson" also appears — no issue
johnson_issues = [i for i in issues if "Johnson" in i.description or "johnson" in i.description.lower()]
assert len(johnson_issues) == 0
def test_last_name_without_full_name(self, tmp_corpus, sample_manifest, sample_fact_registry):
file_path = tmp_corpus / "data" / "emails" / "test.md"
file_path.write_text("Williams reviewed the contract and approved it.")
issues = _check_name_consistency(tmp_corpus, [sample_manifest[0]], sample_fact_registry)
# "Williams" appears but "Bob Williams" does not — should warn
williams_issues = [i for i in issues if "williams" in i.description.lower()]
assert len(williams_issues) == 1
assert williams_issues[0].severity == "warning"
def test_no_people_in_registry(self, tmp_corpus, sample_manifest):
registry_no_people = {"people": [], "organizations": []}
issues = _check_name_consistency(tmp_corpus, sample_manifest, registry_no_people)
assert len(issues) == 0
# ---------------------------------------------------------------------------
# Tests: _check_cross_references
# ---------------------------------------------------------------------------
class TestCheckCrossReferences:
def test_valid_cross_references(self, sample_manifest):
issues = _check_cross_references(Path("/unused"), sample_manifest)
# f001 refs f002, f002 refs f001 — both valid
assert len(issues) == 0
def test_broken_cross_reference(self):
manifest = [
{
"file_id": "f001",
"path": "data/test.md",
"cross_references": ["f999"],
},
]
issues = _check_cross_references(Path("/unused"), manifest)
assert len(issues) == 1
assert issues[0].issue_type == "cross_ref_broken"
assert issues[0].severity == "error"
def test_cross_reference_by_path(self, sample_manifest):
# Add a cross-reference by path instead of file_id
manifest = [
{
"file_id": "f001",
"path": "data/emails/test.md",
"cross_references": ["data/contracts/contract.md"],
},
{
"file_id": "f002",
"path": "data/contracts/contract.md",
"cross_references": [],
},
]
issues = _check_cross_references(Path("/unused"), manifest)
assert len(issues) == 0
def test_empty_cross_references(self):
manifest = [
{"file_id": "f001", "path": "data/test.md", "cross_references": []},
]
issues = _check_cross_references(Path("/unused"), manifest)
assert len(issues) == 0
# ---------------------------------------------------------------------------
# Tests: validate_corpus (integration)
# ---------------------------------------------------------------------------
class TestValidateCorpus:
def test_full_validation(self, tmp_corpus, sample_manifest, sample_fact_registry):
"""Integration test: run full validation on a minimal corpus."""
# Write content with all required facts for f001
file_path = tmp_corpus / "data" / "emails" / "test.md"
content = (
"Alice Johnson confirmed the Q1 marketing budget of $2,034.50. "
"The booking reference is BK-2026-0422. "
"The deadline is April 22, 2026. "
)
# Pad to meet token minimum
content += "Additional context and details follow. " * 800
file_path.write_text(content)
report = asyncio.run(validate_corpus(tmp_corpus, sample_manifest, sample_fact_registry))
assert report.total_files == 2
assert report.files_checked == 1 # only f001 exists
# Should have at least a file_missing error for f002
missing_errors = [i for i in report.errors if i.issue_type == "file_missing"]
assert len(missing_errors) >= 1
# Report should be written to disk
report_path = tmp_corpus / "validation_report.json"
assert report_path.exists()
report_data = json.loads(report_path.read_text())
assert report_data["total_files"] == 2
# ---------------------------------------------------------------------------
# Tests: repair_files
# ---------------------------------------------------------------------------
class TestRepairFiles:
def test_no_errors_skips_repair(self, tmp_corpus, sample_manifest, sample_fact_registry):
"""If no errors, repair_files returns the same report."""
report = ValidationReport(total_files=2, files_checked=2, issues=[])
result = asyncio.run(
repair_files(tmp_corpus, report, sample_manifest, sample_fact_registry)
)
# No errors → re-validates, returns a new report
assert result.total_files == 2
@patch("validator.llm_call", new_callable=AsyncMock)
def test_repair_calls_llm(self, mock_llm, tmp_corpus, sample_manifest, sample_fact_registry):
"""Repair should call the LLM for files with errors."""
# Create a report with one error
issues = [
ValidationIssue(
file_id="f001",
issue_type="token_count",
severity="error",
description="File has 100 tokens, below minimum 4000",
details={"tokens": 100, "min": 4000, "path": "data/emails/test.md"},
),
]
report = ValidationReport(total_files=2, files_checked=1, issues=issues)
# Mock LLM to return padded content
mock_llm.return_value = "Repaired content. " * 1000
result = asyncio.run(
repair_files(tmp_corpus, report, sample_manifest, sample_fact_registry)
)
# LLM should have been called once for f001
assert mock_llm.call_count >= 1
# The repaired file should exist
assert (tmp_corpus / "data" / "emails" / "test.md").exists()

File diff suppressed because it is too large Load diff

305
data-generator/utils.py Normal file
View file

@ -0,0 +1,305 @@
"""Shared utilities: LLM client wrapper, token counting, retry logic, file I/O."""
from __future__ import annotations
import asyncio
import json
import logging
import os
import time
from pathlib import Path
from typing import Any
import litellm
import tiktoken
logger = logging.getLogger(__name__)
# Suppress litellm noise
litellm.suppress_debug_info = True
logging.getLogger("LiteLLM").setLevel(logging.WARNING)
logging.getLogger("litellm").setLevel(logging.WARNING)
# ---------------------------------------------------------------------------
# Token counting
# ---------------------------------------------------------------------------
_enc: tiktoken.Encoding | None = None
def _get_encoder() -> tiktoken.Encoding:
global _enc
if _enc is None:
_enc = tiktoken.get_encoding("cl100k_base")
return _enc
def count_tokens(text: str) -> int:
"""Count tokens using cl100k_base (GPT-4 / Claude approximate)."""
return len(_get_encoder().encode(text))
def estimate_chars_for_tokens(target_tokens: int) -> int:
"""Rough estimate: 1 token ~ 4 characters for English prose."""
return target_tokens * 4
# ---------------------------------------------------------------------------
# LLM client
# ---------------------------------------------------------------------------
DEFAULT_MODEL = "gemini/gemini-2.5-pro"
FAST_MODEL = "gemini/gemini-2.5-flash"
# Rate limiting
_semaphore: asyncio.Semaphore | None = None
def get_semaphore(max_concurrent: int = 10) -> asyncio.Semaphore:
global _semaphore
if _semaphore is None or _semaphore._value != max_concurrent:
_semaphore = asyncio.Semaphore(max_concurrent)
return _semaphore
def _default_temperature(model: str) -> float:
"""Return a sensible default temperature per model family.
Gemini uses 1.0 as its "normal" temperature.
Anthropic/OpenAI treat 1.0 as quite high 0.7 is a better default for creative
prose, and 0.1 for structured/JSON output.
"""
model_lower = model.lower()
if "gemini" in model_lower:
return 1.0
return 0.7
async def llm_call(
prompt: str,
*,
model: str = DEFAULT_MODEL,
system: str | None = None,
temperature: float | None = None,
max_tokens: int = 16384,
json_mode: bool = False,
max_retries: int = 3,
retry_delay: float = 5.0,
max_concurrent: int = 10,
) -> str:
"""Make an LLM call with retry logic and rate limiting.
Args:
temperature: If None, uses a model-aware default (1.0 for Gemini, 0.7 for others).
Returns the raw text response.
"""
if temperature is None:
temperature = _default_temperature(model)
sem = get_semaphore(max_concurrent)
messages: list[dict[str, str]] = []
if system:
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": prompt})
kwargs: dict[str, Any] = {
"model": model,
"messages": messages,
"max_tokens": max_tokens,
"temperature": temperature,
"timeout": 300,
}
# Gemini has limited JSON schema support through litellm — we always parse
# JSON from the raw text response instead of relying on response_format.
# Only enable json_mode for providers that support it reliably.
if json_mode and "gemini" not in model.lower():
kwargs["response_format"] = {"type": "json_object"}
last_error: Exception | None = None
for attempt in range(1, max_retries + 1):
async with sem:
try:
t0 = time.monotonic()
response = await litellm.acompletion(**kwargs)
elapsed = time.monotonic() - t0
text = response.choices[0].message.content
if not text:
logger.warning(f"Empty response from {model} (attempt {attempt})")
continue
input_tokens = getattr(response.usage, "prompt_tokens", 0)
output_tokens = getattr(response.usage, "completion_tokens", 0)
logger.debug(
f"LLM call: model={model} attempt={attempt} "
f"elapsed={elapsed:.1f}s in={input_tokens} out={output_tokens}"
)
return text
except Exception as e:
last_error = e
logger.warning(f"LLM call failed (attempt {attempt}/{max_retries}): {e}")
if attempt < max_retries:
await asyncio.sleep(retry_delay * attempt)
raise RuntimeError(f"All {max_retries} LLM attempts failed. Last error: {last_error}")
async def llm_call_json(
prompt: str,
*,
model: str = DEFAULT_MODEL,
system: str | None = None,
temperature: float | None = None,
max_tokens: int = 16384,
max_retries: int = 3,
max_concurrent: int = 10,
) -> dict[str, Any]:
"""Make an LLM call and parse the response as JSON.
For Gemini models (which have limited JSON schema support), we append an
explicit instruction to return JSON and parse from the raw text response.
For other providers, we use response_format=json_object.
"""
# For Gemini, add explicit JSON instruction since we can't rely on json_mode
effective_prompt = prompt
if "gemini" in model.lower() and "json" not in prompt.lower()[-200:]:
effective_prompt = prompt + "\n\nIMPORTANT: Return your response as a single valid JSON object. No markdown, no explanation — just the JSON."
text = await llm_call(
effective_prompt,
model=model,
system=system,
temperature=temperature,
max_tokens=max_tokens,
json_mode=True,
max_retries=max_retries,
max_concurrent=max_concurrent,
)
return parse_json_response(text)
def parse_json_response(text: str) -> dict[str, Any]:
"""Parse LLM response as JSON, handling code blocks and partial JSON."""
cleaned = text.strip()
# Strip markdown code blocks
if cleaned.startswith("```"):
lines = cleaned.split("\n")
if lines[0].startswith("```"):
lines = lines[1:]
if lines and lines[-1].strip() == "```":
lines = lines[:-1]
cleaned = "\n".join(lines)
try:
return json.loads(cleaned)
except json.JSONDecodeError:
# Try to find JSON object or array
for start_char, end_char in [("{", "}"), ("[", "]")]:
start = cleaned.find(start_char)
end = cleaned.rfind(end_char) + 1
if start >= 0 and end > start:
try:
return json.loads(cleaned[start:end])
except json.JSONDecodeError:
continue
raise ValueError(f"Could not parse JSON from LLM response: {cleaned[:200]}...")
# ---------------------------------------------------------------------------
# File I/O helpers
# ---------------------------------------------------------------------------
def write_json(path: Path, data: Any) -> None:
"""Write JSON atomically."""
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_suffix(".tmp")
with open(tmp, "w") as f:
json.dump(data, f, indent=2, ensure_ascii=False)
tmp.rename(path)
def read_json(path: Path) -> Any:
"""Read JSON file."""
with open(path) as f:
return json.load(f)
def write_text(path: Path, text: str) -> None:
"""Write text file atomically."""
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_suffix(".tmp")
with open(tmp, "w") as f:
f.write(text)
tmp.rename(path)
def read_text(path: Path) -> str:
"""Read text file."""
with open(path) as f:
return f.read()
# ---------------------------------------------------------------------------
# Generation log
# ---------------------------------------------------------------------------
class GenerationLog:
"""Tracks generation progress and stats for checkpointing/resume."""
def __init__(self, log_path: Path):
self.log_path = log_path
self.entries: dict[str, dict[str, Any]] = {}
if log_path.exists():
self.entries = read_json(log_path)
def log_file(
self,
file_id: str,
*,
model: str,
tokens_in: int = 0,
tokens_out: int = 0,
retries: int = 0,
status: str = "ok",
error: str | None = None,
elapsed_s: float = 0.0,
) -> None:
self.entries[file_id] = {
"model": model,
"tokens_in": tokens_in,
"tokens_out": tokens_out,
"retries": retries,
"status": status,
"error": error,
"elapsed_s": round(elapsed_s, 2),
"timestamp": time.time(),
}
self.save()
def is_done(self, file_id: str) -> bool:
entry = self.entries.get(file_id)
return entry is not None and entry.get("status") == "ok"
def save(self) -> None:
write_json(self.log_path, self.entries)
def summary(self) -> dict[str, Any]:
total = len(self.entries)
ok = sum(1 for e in self.entries.values() if e.get("status") == "ok")
failed = sum(1 for e in self.entries.values() if e.get("status") == "failed")
total_tokens_in = sum(e.get("tokens_in", 0) for e in self.entries.values())
total_tokens_out = sum(e.get("tokens_out", 0) for e in self.entries.values())
return {
"total_files": total,
"ok": ok,
"failed": failed,
"total_tokens_in": total_tokens_in,
"total_tokens_out": total_tokens_out,
}

630
data-generator/validator.py Normal file
View file

@ -0,0 +1,630 @@
"""Phase 6: Cross-Reference & Consistency Audit.
After all files are generated, this module audits the corpus for consistency,
checking token counts, locked fact presence, name consistency, and cross-reference
integrity.
"""
from __future__ import annotations
import logging
import re
import statistics
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from utils import FAST_MODEL, count_tokens, llm_call, read_json, read_text, write_json
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Data structures
# ---------------------------------------------------------------------------
@dataclass
class ValidationIssue:
file_id: str
issue_type: str # "token_count" | "missing_fact" | "name_inconsistency" | "date_inconsistency" | "cross_ref_broken"
severity: str # "error" | "warning"
description: str
details: dict = field(default_factory=dict)
@dataclass
class ValidationReport:
total_files: int
files_checked: int
issues: list[ValidationIssue] = field(default_factory=list)
token_stats: dict = field(default_factory=dict) # min, max, mean, median token counts
@property
def errors(self) -> list[ValidationIssue]:
return [i for i in self.issues if i.severity == "error"]
@property
def warnings(self) -> list[ValidationIssue]:
return [i for i in self.issues if i.severity == "warning"]
# ---------------------------------------------------------------------------
# Token-count bounds (slightly relaxed from 5000-10000 to allow minor variance)
# ---------------------------------------------------------------------------
TOKEN_MIN = 4000
TOKEN_MAX = 10500
# ---------------------------------------------------------------------------
# Internal check helpers
# ---------------------------------------------------------------------------
def _check_file_existence(
output_dir: Path,
manifest: list[dict],
) -> list[ValidationIssue]:
"""Check that every file in the manifest exists on disk."""
issues: list[ValidationIssue] = []
for entry in manifest:
file_id = entry.get("file_id", "unknown")
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
if not full_path.exists():
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="file_missing",
severity="error",
description=f"File not found on disk: {rel_path}",
details={"expected_path": str(full_path)},
)
)
return issues
def _check_token_counts(
output_dir: Path,
manifest: list[dict],
) -> tuple[list[ValidationIssue], dict[str, int], dict]:
"""Check token counts for every file. Returns (issues, token_map, token_stats)."""
issues: list[ValidationIssue] = []
token_map: dict[str, int] = {} # file_id -> token count
for entry in manifest:
file_id = entry.get("file_id", "unknown")
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
if not full_path.exists():
continue # already reported by _check_file_existence
content = read_text(full_path)
tokens = count_tokens(content)
token_map[file_id] = tokens
if tokens < TOKEN_MIN:
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="token_count",
severity="error",
description=f"File has {tokens} tokens, below minimum {TOKEN_MIN}",
details={"tokens": tokens, "min": TOKEN_MIN, "path": rel_path},
)
)
elif tokens > TOKEN_MAX:
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="token_count",
severity="error",
description=f"File has {tokens} tokens, above maximum {TOKEN_MAX}",
details={"tokens": tokens, "max": TOKEN_MAX, "path": rel_path},
)
)
# Compute stats
counts = list(token_map.values())
token_stats: dict[str, Any] = {}
if counts:
token_stats = {
"min": min(counts),
"max": max(counts),
"mean": round(statistics.mean(counts), 1),
"median": round(statistics.median(counts), 1),
"total_files_measured": len(counts),
}
return issues, token_map, token_stats
def _normalize_date(date_str: str) -> list[str]:
"""Generate variant string forms of a date for fuzzy matching.
Given "2026-04-22", returns variants like:
- "2026-04-22"
- "April 22, 2026"
- "Apr 22, 2026"
- "04/22/2026"
- "22 April 2026"
"""
import calendar
variants: list[str] = [date_str]
match = re.match(r"(\d{4})-(\d{2})-(\d{2})", date_str)
if match:
year, month_s, day_s = match.groups()
month = int(month_s)
day = int(day_s)
if 1 <= month <= 12:
month_full = calendar.month_name[month]
month_abbr = calendar.month_abbr[month]
# "April 22, 2026"
variants.append(f"{month_full} {day}, {year}")
# "Apr 22, 2026"
variants.append(f"{month_abbr} {day}, {year}")
# "04/22/2026"
variants.append(f"{month_s}/{day_s}/{year}")
# "22 April 2026"
variants.append(f"{day} {month_full} {year}")
# Without leading zero: "4/22/2026"
variants.append(f"{month}/{day_s}/{year}")
# "April 22 2026" (no comma)
variants.append(f"{month_full} {day} {year}")
return variants
def _check_locked_facts(
output_dir: Path,
manifest: list[dict],
fact_registry: dict,
) -> list[ValidationIssue]:
"""Check that locked facts appear in the files that reference them.
Uses pragmatic string matching:
- Dollar amounts: check the dollar string appears (e.g. "$2,034")
- Dates: check any common date format variant appears
- Names: check full name appears at least once
- Reference codes: exact string match
"""
issues: list[ValidationIssue] = []
# Build a lookup: fact_id -> fact dict
fact_lookup: dict[str, dict] = {}
for category in ("financial", "references", "dates", "locations", "domain_facts"):
for fact in fact_registry.get(category, []):
fid = fact.get("id", "")
if fid:
fact_lookup[fid] = {**fact, "_category": category}
# Also index people by id
for person in fact_registry.get("people", []):
pid = person.get("id", "")
if pid:
fact_lookup[pid] = {**person, "_category": "people"}
# Also index organizations by id
for org in fact_registry.get("organizations", []):
oid = org.get("id", "")
if oid:
fact_lookup[oid] = {**org, "_category": "organizations"}
for entry in manifest:
file_id = entry.get("file_id", "unknown")
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
locked_facts = entry.get("locked_facts", [])
if not full_path.exists() or not locked_facts:
continue
content = read_text(full_path)
content_lower = content.lower()
for fact_id in locked_facts:
fact = fact_lookup.get(fact_id)
if fact is None:
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="missing_fact",
severity="warning",
description=f"Locked fact '{fact_id}' not found in fact registry",
details={"fact_id": fact_id},
)
)
continue
category = fact.get("_category", "")
found = False
if category == "financial":
# Check the dollar amount string appears
value = fact.get("value", "")
if value and value in content:
found = True
elif category == "references":
# Exact string match for reference codes
value = fact.get("value", "")
if value and value in content:
found = True
elif category == "dates":
# Check any date format variant appears
date_str = fact.get("date", "")
if date_str:
variants = _normalize_date(date_str)
for variant in variants:
if variant.lower() in content_lower:
found = True
break
elif category == "people":
# Check the full name appears at least once
full_name = fact.get("full_name", "")
if full_name and full_name.lower() in content_lower:
found = True
elif category == "organizations":
# Check the org name appears
name = fact.get("name", "")
if name and name.lower() in content_lower:
found = True
elif category == "locations":
# Check the location name appears
name = fact.get("name", "")
if name and name.lower() in content_lower:
found = True
elif category == "domain_facts":
# Check the fact string appears (partial match)
fact_text = fact.get("fact", "")
if fact_text:
# Check a significant portion of the fact appears
# Use first 40 chars as a reasonable substring
snippet = fact_text[:40].lower()
if snippet in content_lower:
found = True
else:
# Try individual key terms (words > 5 chars)
words = [w for w in fact_text.split() if len(w) > 5]
if words and all(w.lower() in content_lower for w in words[:3]):
found = True
else:
# Unknown category — skip gracefully
continue
if not found:
fact_desc = fact.get("value") or fact.get("full_name") or fact.get("name") or fact.get("date") or fact.get("fact", "")
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="missing_fact",
severity="error",
description=f"Locked fact '{fact_id}' ({category}) not found in file content",
details={
"fact_id": fact_id,
"category": category,
"expected_value": str(fact_desc)[:200],
"path": rel_path,
},
)
)
return issues
def _check_name_consistency(
output_dir: Path,
manifest: list[dict],
fact_registry: dict,
) -> list[ValidationIssue]:
"""Check that person names from the fact registry are spelled consistently.
Looks for partial name matches that differ from the canonical full_name,
which could indicate an inconsistency (e.g. "John Smith" vs "Jon Smith").
"""
issues: list[ValidationIssue] = []
people = fact_registry.get("people", [])
if not people:
return issues
# Collect all person names
name_map: dict[str, str] = {} # last_name_lower -> canonical full_name
for person in people:
full_name = person.get("full_name", "")
if not full_name:
continue
parts = full_name.strip().split()
if len(parts) >= 2:
last_name = parts[-1].lower()
name_map[last_name] = full_name
# For each file, check that if a last name appears, the full canonical name
# also appears somewhere in the file
for entry in manifest:
file_id = entry.get("file_id", "unknown")
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
if not full_path.exists():
continue
content = read_text(full_path)
content_lower = content.lower()
for last_name_lower, canonical_name in name_map.items():
# Only check if the last name appears in the file
if last_name_lower not in content_lower:
continue
# Check that the canonical full name also appears
if canonical_name.lower() not in content_lower:
# The last name is present but the full canonical name is not.
# This might be intentional (using just a last name in dialogue),
# so make it a warning.
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="name_inconsistency",
severity="warning",
description=(
f"Last name '{last_name_lower}' appears but canonical "
f"full name '{canonical_name}' not found in file"
),
details={
"last_name": last_name_lower,
"canonical_name": canonical_name,
"path": rel_path,
},
)
)
return issues
def _check_cross_references(
output_dir: Path,
manifest: list[dict],
) -> list[ValidationIssue]:
"""Check cross-reference integrity.
For each cross_reference in the manifest, verify that both source and target
files exist in the manifest.
"""
issues: list[ValidationIssue] = []
# Build set of valid file_ids
valid_ids = {entry.get("file_id") for entry in manifest}
# Build set of valid paths
valid_paths = {entry.get("path") for entry in manifest}
for entry in manifest:
file_id = entry.get("file_id", "unknown")
cross_refs = entry.get("cross_references", [])
for ref in cross_refs:
# cross_references can be file_ids or paths
if ref not in valid_ids and ref not in valid_paths:
issues.append(
ValidationIssue(
file_id=file_id,
issue_type="cross_ref_broken",
severity="error",
description=f"Cross-reference '{ref}' does not match any file_id or path in the manifest",
details={"reference": ref, "source_file_id": file_id},
)
)
return issues
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
async def validate_corpus(
output_dir: Path,
manifest: list[dict],
fact_registry: dict,
) -> ValidationReport:
"""Run all validation checks on a generated corpus.
Checks:
1. File existence: every file in the manifest must exist on disk
2. Token count: every file must have 4000-10500 tokens
3. Locked facts: for each file, check that its locked_facts appear in the content
4. Name consistency: person names spelled identically everywhere they appear
5. Cross-reference integrity: both source and target files must exist
"""
total_files = len(manifest)
all_issues: list[ValidationIssue] = []
logger.info("Phase 6: Validating corpus (%d files) …", total_files)
# 1. File existence
existence_issues = _check_file_existence(output_dir, manifest)
all_issues.extend(existence_issues)
# Count files that actually exist for reporting
existing_paths = set()
for entry in manifest:
rel_path = entry.get("path", "")
if (output_dir / rel_path).exists():
existing_paths.add(rel_path)
files_checked = len(existing_paths)
# 2. Token counts
token_issues, token_map, token_stats = _check_token_counts(output_dir, manifest)
all_issues.extend(token_issues)
# 3. Locked facts
fact_issues = _check_locked_facts(output_dir, manifest, fact_registry)
all_issues.extend(fact_issues)
# 4. Name consistency
name_issues = _check_name_consistency(output_dir, manifest, fact_registry)
all_issues.extend(name_issues)
# 5. Cross-reference integrity
xref_issues = _check_cross_references(output_dir, manifest)
all_issues.extend(xref_issues)
report = ValidationReport(
total_files=total_files,
files_checked=files_checked,
issues=all_issues,
token_stats=token_stats,
)
logger.info(
"Phase 6 complete — %d errors, %d warnings (checked %d/%d files)",
len(report.errors),
len(report.warnings),
files_checked,
total_files,
)
# Write report to disk
report_path = output_dir / "validation_report.json"
write_json(
report_path,
{
"total_files": report.total_files,
"files_checked": report.files_checked,
"errors": len(report.errors),
"warnings": len(report.warnings),
"token_stats": report.token_stats,
"issues": [
{
"file_id": i.file_id,
"issue_type": i.issue_type,
"severity": i.severity,
"description": i.description,
"details": i.details,
}
for i in report.issues
],
},
)
return report
async def repair_files(
output_dir: Path,
report: ValidationReport,
manifest: list[dict],
fact_registry: dict,
model: str = FAST_MODEL,
) -> ValidationReport:
"""Attempt to regenerate files that failed validation.
Only repairs files with 'error' severity issues.
Returns a new validation report after repairs.
"""
error_file_ids = {issue.file_id for issue in report.errors}
if not error_file_ids:
logger.info("No errors to repair.")
return report
logger.info("Attempting to repair %d files with errors …", len(error_file_ids))
# Build manifest lookup
manifest_lookup = {entry["file_id"]: entry for entry in manifest}
# Build fact lookup
fact_lookup: dict[str, dict] = {}
for category in ("financial", "references", "dates", "locations", "domain_facts", "people", "organizations"):
for fact in fact_registry.get(category, []):
fid = fact.get("id", "")
if fid:
fact_lookup[fid] = {**fact, "_category": category}
for file_id in error_file_ids:
entry = manifest_lookup.get(file_id)
if entry is None:
logger.warning("Cannot repair %s — not found in manifest", file_id)
continue
rel_path = entry.get("path", "")
full_path = output_dir / rel_path
# Collect the specific issues for this file
file_issues = [i for i in report.errors if i.file_id == file_id]
issue_descriptions = "\n".join(f"- {i.description}" for i in file_issues)
# Read current content if file exists
current_content = ""
if full_path.exists():
current_content = read_text(full_path)
# Build list of locked facts with their values
locked_facts_info = []
for fact_id in entry.get("locked_facts", []):
fact = fact_lookup.get(fact_id)
if fact:
cat = fact.get("_category", "unknown")
val = fact.get("value") or fact.get("full_name") or fact.get("name") or fact.get("date") or fact.get("fact", "")
locked_facts_info.append(f" - {fact_id} ({cat}): {val}")
locked_facts_str = "\n".join(locked_facts_info) if locked_facts_info else " (none)"
target_tokens = entry.get("target_tokens", [5000, 10000])
repair_prompt = f"""You are repairing a generated file that failed validation.
## File Details
- file_id: {file_id}
- path: {rel_path}
- format: {entry.get('format', 'unknown')}
- brief: {entry.get('brief', '')}
- tone: {entry.get('tone', '')}
- target tokens: {target_tokens[0]}-{target_tokens[1]}
## Validation Issues
{issue_descriptions}
## Locked Facts (MUST appear in the output)
{locked_facts_str}
## Current Content
{current_content[:8000] if current_content else '(file does not exist — generate from scratch)'}
## Instructions
Rewrite (or generate) the file content to fix ALL validation issues above.
- Ensure the file is between {target_tokens[0]} and {target_tokens[1]} tokens
- Ensure all locked facts appear in the content with their exact values
- Maintain the specified format and tone
- Output ONLY the file content, nothing else no markdown fences or explanations
"""
try:
repaired_content = await llm_call(
repair_prompt,
model=model,
max_tokens=16384,
)
# Write repaired file
full_path.parent.mkdir(parents=True, exist_ok=True)
with open(full_path, "w") as f:
f.write(repaired_content)
logger.info("Repaired file %s (%s)", file_id, rel_path)
except Exception as e:
logger.error("Failed to repair file %s: %s", file_id, e)
# Re-validate after repairs
logger.info("Re-validating after repairs …")
return await validate_corpus(output_dir, manifest, fact_registry)

552
data-generator/worker.py Normal file
View file

@ -0,0 +1,552 @@
"""Phase 5: Parallel file generation workers.
Takes clusters of file entries, a fact registry shard, and optionally
already-generated files for cross-reference context. Generates each file
sequentially within a cluster, passing previously generated files as context.
Clusters at the same topological level run in parallel.
"""
from __future__ import annotations
import asyncio
import logging
import time
from pathlib import Path
from typing import Any
from prompts.file_gen import format_file_gen_prompt, format_retry_prompt
from utils import (
DEFAULT_MODEL,
GenerationLog,
count_tokens,
llm_call,
read_text,
write_text,
)
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
MAX_CONTEXT_TOKENS_PER_FILE = 3000
MAX_TOTAL_CONTEXT_TOKENS = 15000
MAX_RETRIES = 2
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _get_cluster_file_ids(cluster: Any) -> list[str]:
"""Extract file IDs from a cluster object.
Supports both the ``file_ids`` attribute (list[str]) and the
``file_entries`` attribute (list[dict] with ``file_id`` keys) used by the
Cluster dataclass in ``clusterer.py``.
"""
if hasattr(cluster, "file_ids"):
return list(cluster.file_ids)
if hasattr(cluster, "file_entries"):
return [e.get("file_id", "") for e in cluster.file_entries if e.get("file_id")]
raise TypeError(
f"Cluster object has neither 'file_ids' nor 'file_entries': {type(cluster)}"
)
def _truncate_to_tokens(text: str, max_tokens: int) -> str:
"""Truncate *text* to approximately *max_tokens* tokens.
Uses a character-based heuristic first (1 token 4 chars) for speed, then
verifies with the real tokeniser and trims further if needed.
"""
# Fast character-based pre-filter
approx_chars = max_tokens * 4
if len(text) <= approx_chars:
# Likely already within budget verify
if count_tokens(text) <= max_tokens:
return text
# Trim to approximate char limit, then refine
trimmed = text[:approx_chars]
while count_tokens(trimmed) > max_tokens and len(trimmed) > 200:
# Remove ~10% each iteration
trimmed = trimmed[: int(len(trimmed) * 0.9)]
return trimmed
def _build_context_files(
file_entry: dict,
generated: dict[str, str],
manifest_entries: dict[str, dict],
) -> dict[str, str]:
"""Build the context dict for cross-referenced files.
Priority:
- If the referenced file has been generated: include its content (up to
``MAX_CONTEXT_TOKENS_PER_FILE`` tokens).
- If not yet generated: the brief from the manifest is used later by the
prompt builder (via ``manifest_entries``), so we don't duplicate it here.
If total context would exceed ``MAX_TOTAL_CONTEXT_TOKENS``, we keep the
files with the most cross-references first and drop the rest.
"""
cross_refs: list[str] = file_entry.get("cross_references", [])
if not cross_refs:
return {}
# Collect available generated content for referenced files
candidates: list[tuple[str, str]] = []
for ref_id in cross_refs:
content = generated.get(ref_id)
if content is not None:
truncated = _truncate_to_tokens(content, MAX_CONTEXT_TOKENS_PER_FILE)
candidates.append((ref_id, truncated))
# Sort by number of cross-references each candidate has (most connected first)
def _xref_count(file_id: str) -> int:
entry = manifest_entries.get(file_id, {})
return len(entry.get("cross_references", []))
candidates.sort(key=lambda pair: _xref_count(pair[0]), reverse=True)
# Enforce total context budget
context: dict[str, str] = {}
total_tokens = 0
for fid, content in candidates:
tok = count_tokens(content)
if total_tokens + tok > MAX_TOTAL_CONTEXT_TOKENS:
# Try to fit a smaller portion
remaining = MAX_TOTAL_CONTEXT_TOKENS - total_tokens
if remaining > 200:
content = _truncate_to_tokens(content, remaining)
context[fid] = content
break
context[fid] = content
total_tokens += tok
return context
def _validate_content(
content: str,
file_entry: dict,
fact_shard: dict,
) -> list[str]:
"""Validate generated content. Returns a list of issue descriptions (empty = valid)."""
issues: list[str] = []
token_count = count_tokens(content)
# --- Token range check ---
target_tokens = file_entry.get("target_tokens", [5000, 10000])
target_min = target_tokens[0] if isinstance(target_tokens, list) else 5000
target_max = target_tokens[1] if isinstance(target_tokens, list) else 10000
if token_count < target_min:
issues.append(
f"Too short: {token_count:,} tokens (minimum {target_min:,}). "
f"Add more realistic content, filler, and noise."
)
elif token_count > target_max * 1.3:
# Allow 30% overshoot before flagging — slight overshoot is better than
# being too short.
issues.append(
f"Too long: {token_count:,} tokens (maximum ~{target_max:,}). "
f"Trim some filler while keeping all locked facts."
)
# --- Locked facts spot-check ---
locked_ids = set(file_entry.get("locked_facts", []))
if locked_ids:
content_lower = content.lower()
missing_facts: list[str] = []
for category in ("financial", "dates", "references", "locations", "domain_facts"):
for fact in fact_shard.get(category, []):
if fact.get("id") not in locked_ids:
continue
# Determine key values to check in the content
key_values = _extract_key_values(fact, category)
found_any = any(
kv.lower() in content_lower for kv in key_values if kv
)
if not found_any and key_values:
missing_facts.append(
f"{fact['id']} (expected one of: {key_values})"
)
if missing_facts:
issues.append(
"Missing locked facts — the following facts were not found "
"in the generated content:\n "
+ "\n ".join(missing_facts)
)
return issues
def _extract_key_values(fact: dict, category: str) -> list[str]:
"""Extract the key string values from a fact that should appear in the document."""
values: list[str] = []
if category == "financial":
val = fact.get("value", "")
if val:
values.append(val)
elif category == "dates":
date_val = fact.get("date", "")
if date_val:
values.append(date_val)
time_val = fact.get("time", "")
if time_val:
values.append(time_val)
elif category == "references":
val = fact.get("value", "")
if val:
values.append(val)
elif category == "locations":
name = fact.get("name", "")
if name:
values.append(name)
addr = fact.get("address", "")
if addr:
values.append(addr)
elif category == "domain_facts":
fact_text = fact.get("fact", "")
if fact_text:
# For domain facts, check for the first significant clause
# (whole fact string may be too long to match literally)
values.append(fact_text)
return values
# ---------------------------------------------------------------------------
# Core generation
# ---------------------------------------------------------------------------
async def generate_file(
file_entry: dict,
fact_shard: dict,
context_files: dict[str, str],
output_dir: Path,
model: str = DEFAULT_MODEL,
gen_log: GenerationLog | None = None,
manifest_entries: dict[str, dict] | None = None,
) -> str:
"""Generate a single file. Returns the generated content.
Args:
file_entry: manifest entry for this file.
fact_shard: relevant portion of fact registry.
context_files: already-generated files this file cross-references
(file_id -> content).
output_dir: base output directory. File is written to
``output_dir / data / <path>``.
model: LLM model to use.
gen_log: optional generation log for tracking.
manifest_entries: full file_id -> manifest entry map (used for
cross-reference briefs of not-yet-generated files).
"""
file_id: str = file_entry.get("file_id", "unknown")
file_path_rel: str = file_entry.get("path", f"{file_id}.md")
dest = output_dir / "data" / file_path_rel
# --- Resume support ---
if gen_log and gen_log.is_done(file_id):
logger.info("Skipping %s — already done (gen_log)", file_id)
if dest.exists():
return read_text(dest)
# Log says done but file missing — regenerate
logger.warning("%s marked done but file missing, regenerating", file_id)
if dest.exists() and gen_log is None:
logger.info("Skipping %s — file exists on disk", file_id)
return read_text(dest)
# --- Build prompt ---
system_prompt, user_prompt = format_file_gen_prompt(
file_entry=file_entry,
fact_shard=fact_shard,
context_files=context_files,
manifest_entries=manifest_entries or {},
)
# --- Generate with retries ---
content: str = ""
last_issues: list[str] = []
retries_used = 0
t0 = time.monotonic()
for attempt in range(1 + MAX_RETRIES):
try:
if attempt == 0:
content = await llm_call(
user_prompt,
system=system_prompt,
model=model,
max_tokens=16384,
)
else:
# Retry with feedback
retry_prompt = format_retry_prompt(
issues=last_issues,
previous_content=content,
original_prompt=user_prompt,
)
content = await llm_call(
retry_prompt,
system=system_prompt,
model=model,
max_tokens=16384,
)
retries_used = attempt
# Strip any markdown code fences the LLM might have wrapped around output
content = _strip_wrapping_fences(content)
# Validate
last_issues = _validate_content(content, file_entry, fact_shard)
if not last_issues:
break
logger.warning(
"%s attempt %d validation issues: %s",
file_id,
attempt + 1,
last_issues,
)
except Exception as exc:
logger.error("%s attempt %d error: %s", file_id, attempt + 1, exc)
last_issues = [f"Generation error: {exc}"]
if attempt == MAX_RETRIES:
# All retries exhausted — log failure and return whatever we have
elapsed = time.monotonic() - t0
if gen_log:
gen_log.log_file(
file_id,
model=model,
retries=retries_used,
status="failed",
error=str(exc),
elapsed_s=elapsed,
)
logger.error(
"Failed to generate %s after %d attempts: %s",
file_id,
MAX_RETRIES + 1,
exc,
)
return content
elapsed = time.monotonic() - t0
# Even if there are remaining issues after all retries, write the best attempt
status = "ok" if not last_issues else "partial"
if last_issues:
logger.warning(
"%s: writing with unresolved issues after %d retries: %s",
file_id,
retries_used,
last_issues,
)
# Write to disk
write_text(dest, content)
logger.info(
"Generated %s (%d tokens, %d retries, %.1fs) -> %s",
file_id,
count_tokens(content),
retries_used,
elapsed,
dest,
)
# Log
if gen_log:
gen_log.log_file(
file_id,
model=model,
tokens_out=count_tokens(content),
retries=retries_used,
status=status,
error="; ".join(last_issues) if last_issues else None,
elapsed_s=elapsed,
)
return content
def _strip_wrapping_fences(text: str) -> str:
"""Remove markdown code fences that an LLM might wrap around the output."""
stripped = text.strip()
if stripped.startswith("```"):
lines = stripped.split("\n")
# Remove opening fence (e.g. ```markdown, ```text, ```)
if lines[0].startswith("```"):
lines = lines[1:]
# Remove closing fence
if lines and lines[-1].strip() == "```":
lines = lines[:-1]
return "\n".join(lines)
return text
# ---------------------------------------------------------------------------
# Cluster-level generation
# ---------------------------------------------------------------------------
async def generate_cluster(
cluster: Any,
manifest_entries: dict[str, dict],
fact_shard: dict,
output_dir: Path,
context_files: dict[str, str],
model: str = DEFAULT_MODEL,
gen_log: GenerationLog | None = None,
) -> dict[str, str]:
"""Generate all files in a cluster sequentially.
Returns dict of file_id -> content for all generated files.
Each file in the cluster sees previously generated files as context.
Args:
cluster: a Cluster object with ``file_entries`` (list[dict]) and
``level`` (int). Each entry dict must contain a ``file_id`` key.
manifest_entries: file_id -> manifest entry for ALL files.
fact_shard: the fact registry (or relevant shard).
output_dir: base output directory.
context_files: files from dependency clusters (file_id -> content).
model: LLM model to use.
gen_log: optional generation log.
"""
# Merge dependency context with what we generate in this cluster
combined_context: dict[str, str] = dict(context_files)
generated: dict[str, str] = {}
# Extract ordered file IDs from the cluster's file_entries list
file_ids = _get_cluster_file_ids(cluster)
for file_id in file_ids:
entry = manifest_entries.get(file_id)
if entry is None:
logger.warning(
"File %s in cluster but not in manifest — skipping", file_id
)
continue
# Build cross-reference context for this specific file
file_context = _build_context_files(entry, combined_context, manifest_entries)
content = await generate_file(
file_entry=entry,
fact_shard=fact_shard,
context_files=file_context,
output_dir=output_dir,
model=model,
gen_log=gen_log,
manifest_entries=manifest_entries,
)
generated[file_id] = content
combined_context[file_id] = content
return generated
# ---------------------------------------------------------------------------
# Top-level orchestrator
# ---------------------------------------------------------------------------
async def generate_all(
clusters: list[Any],
manifest_entries: dict[str, dict],
output_dir: Path,
model: str = DEFAULT_MODEL,
max_concurrent: int = 10,
gen_log: GenerationLog | None = None,
fallback_fact_registry: dict | None = None,
) -> None:
"""Generate all files across all clusters, respecting topological order.
Clusters at the same ``level`` run in parallel (up to *max_concurrent*).
Clusters at different levels run sequentially (lower levels first).
Each cluster uses its own ``cluster.fact_shard`` (set by the clusterer's
sharding logic). If a cluster has no ``fact_shard`` attribute or it is
empty, *fallback_fact_registry* is used instead.
Args:
clusters: list of Cluster-like objects, **already ordered by level**.
Each should have a ``fact_shard`` attribute (dict) set by the
clusterer.
manifest_entries: file_id -> manifest entry for ALL files.
output_dir: base output directory.
model: LLM model to use.
max_concurrent: maximum number of clusters processed in parallel
within a single level.
gen_log: optional generation log.
fallback_fact_registry: full fact registry used when a cluster has no
``fact_shard``.
"""
# Group clusters by level
levels: dict[int, list[Any]] = {}
for cluster in clusters:
level = getattr(cluster, "level", 0)
levels.setdefault(level, []).append(cluster)
# All generated content so far (shared across levels)
all_generated: dict[str, str] = {}
for level_num in sorted(levels.keys()):
level_clusters = levels[level_num]
logger.info(
"Level %d: processing %d cluster(s) (up to %d concurrent)",
level_num,
len(level_clusters),
max_concurrent,
)
sem = asyncio.Semaphore(max_concurrent)
async def _run_cluster(c: Any) -> dict[str, str]:
async with sem:
# Use the cluster's own sharded fact registry; fall back to
# the full registry if the cluster doesn't have one.
cluster_facts = getattr(c, "fact_shard", None) or {}
if not cluster_facts and fallback_fact_registry:
cluster_facts = fallback_fact_registry
# Snapshot current generated content as context for this cluster
return await generate_cluster(
cluster=c,
manifest_entries=manifest_entries,
fact_shard=cluster_facts,
output_dir=output_dir,
context_files=dict(all_generated),
model=model,
gen_log=gen_log,
)
results = await asyncio.gather(
*(_run_cluster(c) for c in level_clusters),
return_exceptions=True,
)
for i, result in enumerate(results):
if isinstance(result, Exception):
try:
cluster_ids = _get_cluster_file_ids(level_clusters[i])
except Exception:
cluster_ids = [f"<cluster index {i}>"]
logger.error(
"Cluster %s at level %d failed: %s",
cluster_ids,
level_num,
result,
)
else:
all_generated.update(result)
if gen_log:
summary = gen_log.summary()
logger.info("Generation complete. Summary: %s", summary)