mirror of
https://github.com/supermemoryai/supermemory.git
synced 2026-10-08 03:08:21 +00:00
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:
parent
cba994be3f
commit
771be5cef8
12 changed files with 343 additions and 188 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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) ---")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
206
data-generator/test_utils.py
Normal file
206
data-generator/test_utils.py
Normal 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")
|
||||
|
|
@ -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())
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue