diff --git a/data-generator/clusterer.py b/data-generator/clusterer.py index a1127c05..add56b92 100644 --- a/data-generator/clusterer.py +++ b/data-generator/clusterer.py @@ -365,7 +365,7 @@ def _detect_and_merge_cycles( merge_map[cid] = primary cluster_groups[primary].extend(cluster_groups.pop(cid)) logger.info( - f"Merged cyclic clusters {scc_sorted[1:]} into '{primary}'" + "Merged cyclic clusters %s into '%s'", scc_sorted[1:], primary ) if not merge_map: @@ -431,8 +431,9 @@ def _topological_sort_with_levels( 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." + "Topological sort did not visit all clusters. " + "Remaining (possible cycle): %s. Assigning max level.", + missing, ) max_level = max((lvl for _, lvl in result), default=0) + 1 for cid in sorted(missing): @@ -481,7 +482,6 @@ def assign_clusters( # 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 @@ -490,7 +490,6 @@ def assign_clusters( 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) diff --git a/data-generator/generate.py b/data-generator/generate.py index 421f16de..508aa34c 100644 --- a/data-generator/generate.py +++ b/data-generator/generate.py @@ -37,7 +37,7 @@ 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 utils import DEFAULT_MODEL, GenerationLog, read_json, read_text from validator import validate_corpus from worker import generate_all @@ -436,18 +436,7 @@ async def run_pipeline( 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 - ], - }) + # validate_corpus already writes validation_report.json to disk return # --- Questions-only mode --- @@ -510,19 +499,7 @@ async def run_pipeline( 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 - ], - }) + # validate_corpus already writes validation_report.json to disk # Phase 7: Questions logger.info("--- QUESTIONS (Phase 7) ---") diff --git a/data-generator/planner.py b/data-generator/planner.py index 9011907b..3d25dfd0 100644 --- a/data-generator/planner.py +++ b/data-generator/planner.py @@ -20,6 +20,7 @@ from utils import ( llm_call, llm_call_json, read_text, + unwrap_json_list, write_json, write_text, ) @@ -231,7 +232,7 @@ def _validate_manifest(manifest: list[dict[str, Any]]) -> list[dict[str, Any]]: if all_warnings: for w in all_warnings: - logger.warning(f"Manifest validation: {w}") + logger.warning("Manifest validation: %s", w) return manifest @@ -365,23 +366,7 @@ async def _generate_small_manifest( 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)}") + return unwrap_json_list(result, expected_keys=("files", "manifest", "entries")) async def _generate_large_manifest( @@ -469,29 +454,7 @@ async def _generate_large_manifest( 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)}" - ) + entries = unwrap_json_list(result, expected_keys=("files", "manifest", "entries")) all_entries.extend(entries) current_start_id += section_file_count @@ -552,11 +515,10 @@ async def run_planning( Returns (brief, facts, manifest). """ - out = Path(output_dir) - out.mkdir(parents=True, exist_ok=True) + output_dir.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) + brief = await generate_scenario_brief(scenario_block, file_count, output_dir, model) + facts = await extract_fact_registry(brief, output_dir, model) + manifest = await generate_manifest(brief, facts, file_count, output_dir, model) return brief, facts, manifest diff --git a/data-generator/prompts/file_gen.py b/data-generator/prompts/file_gen.py index 64c40ae4..cbbd99af 100644 --- a/data-generator/prompts/file_gen.py +++ b/data-generator/prompts/file_gen.py @@ -1,5 +1,7 @@ """Prompt templates for Phase 5: Individual file generation.""" +from utils import estimate_chars_for_tokens + 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. @@ -341,7 +343,7 @@ def _build_cross_reference_context( 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')}" + f"Brief: {entry.get('brief', 'No brief available')}" ) parts.append(f"### {ref_id} (not yet generated — brief only)\n\n{brief}") else: @@ -383,8 +385,8 @@ def format_file_gen_prompt( 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_min_chars = estimate_chars_for_tokens(target_min_tokens) + target_max_chars = estimate_chars_for_tokens(target_max_tokens) target_length_instructions = ( f"- Target: **{target_min_tokens:,}-{target_max_tokens:,} tokens** " @@ -407,7 +409,7 @@ def format_file_gen_prompt( 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."), + file_summary=file_entry.get("brief", "No brief provided."), target_length_instructions=target_length_instructions, format_instructions=format_instructions, format_notes=format_notes, diff --git a/data-generator/questions.py b/data-generator/questions.py index d92bebb2..d9e195f6 100644 --- a/data-generator/questions.py +++ b/data-generator/questions.py @@ -11,7 +11,16 @@ 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 utils import ( + DEFAULT_MODEL, + count_tokens, + llm_call_json, + read_json, + read_text, + truncate_to_tokens, + unwrap_json_list, + write_json, +) from prompts.questions import QUESTION_GEN_PROMPT, QUESTION_GEN_SYSTEM logger = logging.getLogger(__name__) @@ -33,31 +42,9 @@ VALID_FAMILIES = {"single_hop", "multi_hop", "format_spanning", "edit_then_recal # --------------------------------------------------------------------------- -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) + return truncate_to_tokens(scenario_brief, MAX_SCENARIO_TOKENS, add_suffix=True) def _build_fact_summary(fact_registry: dict) -> str: @@ -157,8 +144,8 @@ def _sample_file_excerpts( if not full_path.exists(): continue - content = read_text(full_path) - excerpt = _truncate_to_tokens(content, MAX_EXCERPT_TOKENS) + content = full_path.read_text(encoding="utf-8") + excerpt = truncate_to_tokens(content, MAX_EXCERPT_TOKENS, add_suffix=True) excerpt_tokens = count_tokens(excerpt) if total_tokens + excerpt_tokens > MAX_TOTAL_EXCERPT_TOKENS: @@ -252,7 +239,7 @@ async def generate_questions( # Resume support: skip if already exists if output_path.exists(): logger.info("Phase 7 skipped — question.json already exists") - return read_text(output_path) + return read_json(output_path) logger.info("Phase 7: Generating eval questions …") @@ -284,26 +271,7 @@ async def generate_questions( ) # 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)}") + questions = unwrap_json_list(result, expected_keys=("questions", "eval_questions", "items")) # Validate and clean up questions = _validate_questions(questions, manifest) diff --git a/data-generator/requirements.txt b/data-generator/requirements.txt index 438b150b..cdbc49d0 100644 --- a/data-generator/requirements.txt +++ b/data-generator/requirements.txt @@ -1,6 +1,2 @@ litellm>=1.50.0 -pydantic>=2.0 -pydantic-settings>=2.0 tiktoken>=0.7.0 -pyyaml>=6.0 -asyncio-pool>=0.7.0 diff --git a/data-generator/test_questions.py b/data-generator/test_questions.py index aa5eadef..9af5ad3f 100644 --- a/data-generator/test_questions.py +++ b/data-generator/test_questions.py @@ -16,10 +16,10 @@ from questions import ( _build_manifest_summary, _build_scenario_summary, _sample_file_excerpts, - _truncate_to_tokens, _validate_questions, generate_questions, ) +from utils import truncate_to_tokens as _truncate_to_tokens # --------------------------------------------------------------------------- @@ -164,7 +164,7 @@ class TestTruncateToTokens: def test_long_text_truncated(self): text = "Word " * 5000 # ~5000 tokens - result = _truncate_to_tokens(text, 100) + result = _truncate_to_tokens(text, 100, add_suffix=True) assert len(result) < len(text) assert result.endswith("[… truncated …]") @@ -469,8 +469,10 @@ class TestGenerateQuestions: ) ) - # Should return the existing content (as string since read_text returns string) - assert "Existing?" in str(result) + # Should return the existing content as a parsed list, not a raw string + assert isinstance(result, list), f"Expected list, got {type(result)}: {result!r}" + assert len(result) == 1 + assert result[0]["prompt"] == "Existing?" # --------------------------------------------------------------------------- diff --git a/data-generator/test_utils.py b/data-generator/test_utils.py new file mode 100644 index 00000000..d3f08f1d --- /dev/null +++ b/data-generator/test_utils.py @@ -0,0 +1,206 @@ +"""Tests for utils.py — shared utilities.""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest + +from utils import ( + count_tokens, + estimate_chars_for_tokens, + get_semaphore, + parse_json_response, + truncate_to_tokens, + unwrap_json_list, +) + + +# --------------------------------------------------------------------------- +# Tests: truncate_to_tokens +# --------------------------------------------------------------------------- + + +class TestTruncateToTokens: + def test_short_text_unchanged(self): + text = "Hello world" + assert truncate_to_tokens(text, 1000) == text + + def test_long_text_truncated(self): + text = "word " * 5000 + result = truncate_to_tokens(text, 100) + assert count_tokens(result) <= 100 + + def test_empty_text(self): + assert truncate_to_tokens("", 100) == "" + + def test_exact_boundary(self): + text = "a " * 50 + tokens = count_tokens(text) + assert truncate_to_tokens(text, tokens) == text + + def test_add_suffix_false_no_suffix(self): + """Default (add_suffix=False) should not append the truncation marker.""" + text = "word " * 5000 + result = truncate_to_tokens(text, 100) + assert "[… truncated …]" not in result + + def test_add_suffix_true_appends_marker(self): + """add_suffix=True should append truncation marker.""" + text = "word " * 5000 + result = truncate_to_tokens(text, 100, add_suffix=True) + assert result.endswith("[… truncated …]") + + def test_add_suffix_no_marker_if_not_truncated(self): + """Short text should not get the suffix even if add_suffix=True.""" + text = "Short text." + result = truncate_to_tokens(text, 1000, add_suffix=True) + assert result == text + + def test_add_suffix_sentence_boundary(self): + """add_suffix path should try to break at a sentence boundary.""" + text = "First sentence. Second sentence. Third sentence. " * 500 + result = truncate_to_tokens(text, 10, add_suffix=True) + assert result.endswith("[… truncated …]") + # Should have broken at a boundary + body = result.replace("\n\n[… truncated …]", "") + assert body.endswith(". ") or body.endswith(".") + + +# --------------------------------------------------------------------------- +# Tests: unwrap_json_list +# --------------------------------------------------------------------------- + + +class TestUnwrapJsonList: + def test_bare_list(self): + data = [{"id": "a"}, {"id": "b"}] + assert unwrap_json_list(data) == data + + def test_dict_with_expected_key(self): + data = {"files": [{"id": "a"}]} + assert unwrap_json_list(data, expected_keys=("files",)) == [{"id": "a"}] + + def test_dict_fallback_to_any_list(self): + data = {"unknown_key": [{"id": "a"}]} + assert unwrap_json_list(data) == [{"id": "a"}] + + def test_dict_no_list_values_raises(self): + data = {"key": "string_value"} + with pytest.raises(ValueError, match="Expected a JSON array"): + unwrap_json_list(data) + + def test_unexpected_type_raises(self): + with pytest.raises(ValueError, match="Unexpected response type"): + unwrap_json_list("a string") + + def test_expected_keys_tried_in_order(self): + data = {"entries": [{"id": "e"}], "files": [{"id": "f"}]} + result = unwrap_json_list(data, expected_keys=("files", "entries")) + assert result == [{"id": "f"}] + + def test_empty_list(self): + assert unwrap_json_list([]) == [] + + def test_expected_keys_skip_non_list(self): + data = {"files": "not a list", "entries": [{"id": "a"}]} + result = unwrap_json_list(data, expected_keys=("files", "entries")) + assert result == [{"id": "a"}] + + +# --------------------------------------------------------------------------- +# Tests: estimate_chars_for_tokens +# --------------------------------------------------------------------------- + + +class TestEstimateCharsForTokens: + def test_basic(self): + assert estimate_chars_for_tokens(100) == 400 + assert estimate_chars_for_tokens(0) == 0 + + def test_negative(self): + assert estimate_chars_for_tokens(-1) == -4 + + +# --------------------------------------------------------------------------- +# Tests: get_semaphore +# --------------------------------------------------------------------------- + + +class TestGetSemaphore: + def test_returns_semaphore(self): + import utils + + # Reset global state + utils._semaphore = None + utils._semaphore_capacity = 0 + sem = get_semaphore(5) + assert isinstance(sem, asyncio.Semaphore) + assert sem._value == 5 + + def test_reuses_semaphore_with_same_capacity(self): + import utils + + utils._semaphore = None + utils._semaphore_capacity = 0 + sem1 = get_semaphore(5) + sem2 = get_semaphore(5) + assert sem1 is sem2 + + def test_recreates_semaphore_with_different_capacity(self): + import utils + + utils._semaphore = None + utils._semaphore_capacity = 0 + sem1 = get_semaphore(5) + sem2 = get_semaphore(10) + assert sem1 is not sem2 + assert sem2._value == 10 + + def test_does_not_recreate_on_acquired_permits(self): + """The semaphore should NOT be recreated when permits are acquired. + + This was the bug: the old code checked ``_value != max_concurrent`` + which triggers after any ``acquire()``. + """ + import utils + + utils._semaphore = None + utils._semaphore_capacity = 0 + + async def _run(): + sem = get_semaphore(3) + await sem.acquire() + # After acquire, _value is 2, but capacity is still 3 + sem2 = get_semaphore(3) + assert sem is sem2, "Semaphore should be reused despite acquired permits" + sem.release() + + asyncio.run(_run()) + + +# --------------------------------------------------------------------------- +# Tests: parse_json_response +# --------------------------------------------------------------------------- + + +class TestParseJsonResponse: + def test_bare_json(self): + assert parse_json_response('{"key": "value"}') == {"key": "value"} + + def test_json_with_markdown_fences(self): + text = "```json\n{\"key\": \"value\"}\n```" + assert parse_json_response(text) == {"key": "value"} + + def test_json_array(self): + assert parse_json_response('[1, 2, 3]') == [1, 2, 3] + + def test_json_embedded_in_text(self): + text = "Here is the result: {\"key\": \"value\"} and some trailing text." + assert parse_json_response(text) == {"key": "value"} + + def test_invalid_json_raises(self): + with pytest.raises(ValueError, match="Could not parse JSON"): + parse_json_response("not json at all") diff --git a/data-generator/test_worker.py b/data-generator/test_worker.py index da4a8c10..7f68b94c 100644 --- a/data-generator/test_worker.py +++ b/data-generator/test_worker.py @@ -19,12 +19,12 @@ from worker import ( _extract_key_values, _get_cluster_file_ids, _strip_wrapping_fences, - _truncate_to_tokens, _validate_content, generate_all, generate_cluster, generate_file, ) +from utils import truncate_to_tokens as _truncate_to_tokens # --------------------------------------------------------------------------- # Helpers & fixtures @@ -66,13 +66,13 @@ def _make_entry( ) -> dict: return { "file_id": file_id, - "path": path or f"docs/{file_id}.md", + "path": path or f"data/docs/{file_id}.md", "format": fmt, "date": "2024-03-15", "authors": authors or ["alice"], "author": authors or ["alice"], "tone": "casual", - "summary": f"Test document {file_id}", + "brief": f"Test document {file_id}", "cross_references": cross_refs or [], "locked_facts": locked_facts or [], "target_tokens": target_tokens or [5000, 10000], @@ -414,7 +414,7 @@ class TestGetClusterFileIds: class BadCluster: level = 0 - with pytest.raises(TypeError, match="neither"): + with pytest.raises(TypeError, match="file_entries"): _get_cluster_file_ids(BadCluster()) diff --git a/data-generator/utils.py b/data-generator/utils.py index 89574d9b..6103c052 100644 --- a/data-generator/utils.py +++ b/data-generator/utils.py @@ -5,7 +5,6 @@ from __future__ import annotations import asyncio import json import logging -import os import time from pathlib import Path from typing import Any @@ -44,6 +43,78 @@ def estimate_chars_for_tokens(target_tokens: int) -> int: return target_tokens * 4 +def truncate_to_tokens(text: str, max_tokens: int, *, add_suffix: bool = False) -> str: + """Truncate *text* to approximately *max_tokens* tokens. + + Uses a character-based heuristic first (via :func:`estimate_chars_for_tokens`) + for speed, then verifies with the real tokeniser and trims further if needed. + + Args: + text: The text to truncate. + max_tokens: Maximum number of tokens. + add_suffix: If ``True``, append ``"[… truncated …]"`` when truncation + occurs. Used by the question generator; the worker skips it. + """ + approx_chars = estimate_chars_for_tokens(max_tokens) + if len(text) <= approx_chars and count_tokens(text) <= max_tokens: + return text + + trimmed = text[:approx_chars] + + if add_suffix: + # Try to cut at a sentence/paragraph boundary for cleaner output + for sep in ("\n\n", "\n", ". ", " "): + idx = trimmed.rfind(sep) + if idx > approx_chars // 2: + trimmed = trimmed[: idx + len(sep)] + break + return trimmed + "\n\n[… truncated …]" + + # Iterative refinement for the non-suffix (worker) path + while count_tokens(trimmed) > max_tokens and len(trimmed) > 200: + trimmed = trimmed[: int(len(trimmed) * 0.9)] + return trimmed + + +def unwrap_json_list( + result: Any, + expected_keys: tuple[str, ...] = (), +) -> list[dict]: + """Unwrap an LLM JSON response that should be a list of dicts. + + The LLM may return a bare list or wrap it in a dict with a key like + ``"files"``, ``"manifest"``, ``"entries"``, ``"questions"``, etc. + + Args: + result: Parsed JSON value (list or dict). + expected_keys: Key names to probe in order when *result* is a dict. + After these, any remaining dict values that are lists are tried as + a fallback. + + Returns: + The unwrapped ``list[dict]``. + + Raises: + ValueError: If the result cannot be unwrapped into a list. + """ + if isinstance(result, list): + return result + + if isinstance(result, dict): + for key in expected_keys: + if key in result and isinstance(result[key], list): + return result[key] + # Fallback: first list value we find + for v in result.values(): + if isinstance(v, list): + return v + raise ValueError( + f"Expected a JSON array, got dict with keys: {list(result.keys())}" + ) + + raise ValueError(f"Unexpected response type: {type(result)}") + + # --------------------------------------------------------------------------- # LLM client # --------------------------------------------------------------------------- @@ -53,12 +124,14 @@ FAST_MODEL = "gemini/gemini-2.5-flash" # Rate limiting _semaphore: asyncio.Semaphore | None = None +_semaphore_capacity: int = 0 def get_semaphore(max_concurrent: int = 10) -> asyncio.Semaphore: - global _semaphore - if _semaphore is None or _semaphore._value != max_concurrent: + global _semaphore, _semaphore_capacity + if _semaphore is None or _semaphore_capacity != max_concurrent: _semaphore = asyncio.Semaphore(max_concurrent) + _semaphore_capacity = max_concurrent return _semaphore @@ -127,20 +200,20 @@ async def llm_call( text = response.choices[0].message.content if not text: - logger.warning(f"Empty response from {model} (attempt {attempt})") + logger.warning("Empty response from %s (attempt %d)", model, 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}" + "LLM call: model=%s attempt=%d elapsed=%.1fs in=%d out=%d", + model, attempt, elapsed, input_tokens, output_tokens, ) return text except Exception as e: last_error = e - logger.warning(f"LLM call failed (attempt {attempt}/{max_retries}): {e}") + logger.warning("LLM call failed (attempt %d/%d): %s", attempt, max_retries, e) if attempt < max_retries: await asyncio.sleep(retry_delay * attempt) diff --git a/data-generator/validator.py b/data-generator/validator.py index dec6ddce..eadaca82 100644 --- a/data-generator/validator.py +++ b/data-generator/validator.py @@ -7,6 +7,7 @@ integrity. from __future__ import annotations +import calendar import logging import re import statistics @@ -14,7 +15,7 @@ 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 +from utils import FAST_MODEL, count_tokens, llm_call, read_json, read_text, write_json, write_text logger = logging.getLogger(__name__) @@ -150,8 +151,6 @@ def _normalize_date(date_str: str) -> list[str]: - "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) @@ -616,9 +615,7 @@ Rewrite (or generate) the file content to fix ALL validation issues above. ) # Write repaired file - full_path.parent.mkdir(parents=True, exist_ok=True) - with open(full_path, "w") as f: - f.write(repaired_content) + write_text(full_path, repaired_content) logger.info("Repaired file %s (%s)", file_id, rel_path) diff --git a/data-generator/worker.py b/data-generator/worker.py index 922b0abf..d256e76b 100644 --- a/data-generator/worker.py +++ b/data-generator/worker.py @@ -21,6 +21,7 @@ from utils import ( count_tokens, llm_call, read_text, + truncate_to_tokens, write_text, ) @@ -40,42 +41,14 @@ MAX_RETRIES = 2 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) + """Extract file IDs from a cluster object's ``file_entries`` list.""" 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)}" + f"Cluster object has no 'file_entries' attribute: {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], @@ -101,7 +74,7 @@ def _build_context_files( 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) + 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) @@ -120,7 +93,7 @@ def _build_context_files( # Try to fit a smaller portion remaining = MAX_TOTAL_CONTEXT_TOKENS - total_tokens if remaining > 200: - content = _truncate_to_tokens(content, remaining) + content = truncate_to_tokens(content, remaining) context[fid] = content break context[fid] = content @@ -240,15 +213,15 @@ async def generate_file( context_files: already-generated files this file cross-references (file_id -> content). output_dir: base output directory. File is written to - ``output_dir / data / ``. + ``output_dir / `` (manifest paths already include ``data/``). 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 + file_path_rel: str = file_entry.get("path", f"data/{file_id}.md") + dest = output_dir / file_path_rel # --- Resume support --- if gen_log and gen_log.is_done(file_id):