fix: apply review feedback — fix double data/ prefix, semaphore bug, resume bug, consolidate duplicated code

- Fix worker.py writing to data/data/ instead of data/ (critical path bug)
- Fix semaphore recreation on every call due to checking _value instead of capacity
- Fix questions.py resume returning raw string instead of list[dict]
- Fix prompts/file_gen.py reading 'summary' instead of 'brief' from manifest
- Extract shared unwrap_json_list() and truncate_to_tokens() into utils.py
- Remove redundant validation report writes in generate.py
- Remove unused imports and dependencies
- Fix f-string logger calls to use lazy %s formatting
- Move calendar import to top-level in validator.py
- Use write_text() for atomic writes in repair_files()
- Strengthen test_resume_support to assert return type
This commit is contained in:
Dhravya 2026-04-28 23:49:23 +00:00
parent cba994be3f
commit 771be5cef8
12 changed files with 343 additions and 188 deletions

View file

@ -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)

View file

@ -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) ---")

View file

@ -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

View file

@ -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,

View file

@ -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)

View file

@ -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

View file

@ -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?"
# ---------------------------------------------------------------------------

View file

@ -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")

View file

@ -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())

View file

@ -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)

View file

@ -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)

View file

@ -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 / <path>``.
``output_dir / <path>`` (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):