Skip to content

Commit d1d3f4d

Browse files
fix(compiler): guard concept/entity generation against malformed & truncated LLM output (#161)
* fix(compiler): guard concept/entity generation against malformed & truncated LLM output Two silent-data-loss bugs in _compile_concepts, both from trusting the per-page LLM response shape: - #158: a response returned as a JSON array (e.g. [{...}] or a multi-item list) reached .get() on a list and raised AttributeError, which the local except (JSONDecodeError, ValueError) did not catch — the page was dropped and, because the doc index records it as written, never retried. New _parse_page_json unwraps a single-element [{...}] array (recovering the common case) and returns None for other wrong shapes so the page is skipped cleanly instead of writing the raw JSON as its body. - #148: a response that hit finish_reason=length was repaired by json_repair and written anyway, overwriting an existing concept page with truncated content while still reporting [OK]. _warn_if_truncated now reports whether it truncated, and the four page-generation calls pass raise_on_truncation=True so a truncated response skips the write (existing page preserved). Other callers (plan, summary, overview) keep the warn-only behavior. Closes #158. Closes #148. Claude-Session: https://claude.ai/code/session_01UtbmJxjtw6FtP8fUXUKVtg * refactor(compiler): extract _page_fields + _llm_call_page_async; cover entity/prose paths Addresses code-review findings on the #158/#148 fix, no behavior change: - Collapse the four near-identical parse/fallback blocks in the concept and entity generation closures into a shared _page_fields() helper, so a future edge case is fixed in one place instead of four copies that can drift. - Add _llm_call_page_async() which hard-codes raise_on_truncation=True, so a page-generating call site can't silently forget the truncation guard and regress #148. - Tests: a _page_fields shape unit test (object / single-element-array unwrap / wrong-shape skip / non-JSON prose fallback) and an entity-path truncation-skip integration test. Claude-Session: https://claude.ai/code/session_01UtbmJxjtw6FtP8fUXUKVtg
1 parent dce4972 commit d1d3f4d

2 files changed

Lines changed: 274 additions & 50 deletions

File tree

‎openkb/agent/compiler.py‎

Lines changed: 91 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -392,7 +392,14 @@ def _fmt_messages(messages: list[dict], max_content: int = 200) -> str:
392392
return "\n".join(parts)
393393

394394

395-
def _llm_call(model: str, messages: list[dict], step_name: str, **kwargs) -> str:
395+
class TruncatedResponseError(Exception):
396+
"""Raised when an LLM response hit the length cap and the caller asked to
397+
treat truncation as a failure (so a partial page is skipped, not written)."""
398+
399+
400+
def _llm_call(
401+
model: str, messages: list[dict], step_name: str, raise_on_truncation: bool = False, **kwargs
402+
) -> str:
396403
"""Single LLM call with animated progress and debug logging."""
397404
messages = _prepare_messages(model, messages)
398405
extra_headers = get_extra_headers()
@@ -411,16 +418,22 @@ def _llm_call(model: str, messages: list[dict], step_name: str, **kwargs) -> str
411418

412419
response = litellm.completion(model=model, messages=messages, **kwargs)
413420
content = response.choices[0].message.content or ""
414-
_warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
421+
truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
415422

416423
spinner.stop(_format_usage(time.time() - t0, response.usage))
417424
logger.debug(
418425
"LLM response [%s]:\n%s", step_name, content[:500] + ("..." if len(content) > 500 else "")
419426
)
427+
if raise_on_truncation and truncated:
428+
raise TruncatedResponseError(
429+
f"LLM [{step_name}] hit the length limit; skipping to avoid a truncated page"
430+
)
420431
return content.strip()
421432

422433

423-
async def _llm_call_async(model: str, messages: list[dict], step_name: str, **kwargs) -> str:
434+
async def _llm_call_async(
435+
model: str, messages: list[dict], step_name: str, raise_on_truncation: bool = False, **kwargs
436+
) -> str:
424437
"""Async LLM call with timing output and debug logging."""
425438
messages = _prepare_messages(model, messages)
426439
extra_headers = get_extra_headers()
@@ -437,17 +450,32 @@ async def _llm_call_async(model: str, messages: list[dict], step_name: str, **kw
437450

438451
response = await litellm.acompletion(model=model, messages=messages, **kwargs)
439452
content = response.choices[0].message.content or ""
440-
_warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
453+
truncated = _warn_if_truncated(response, step_name, kwargs.get("max_tokens"))
441454

442455
elapsed = time.time() - t0
443456
sys.stdout.write(f" {step_name}... {_format_usage(elapsed, response.usage)}\n")
444457
sys.stdout.flush()
445458
logger.debug(
446459
"LLM response [%s]:\n%s", step_name, content[:500] + ("..." if len(content) > 500 else "")
447460
)
461+
if raise_on_truncation and truncated:
462+
raise TruncatedResponseError(
463+
f"LLM [{step_name}] hit the length limit; skipping to avoid a truncated page"
464+
)
448465
return content.strip()
449466

450467

468+
async def _llm_call_page_async(model: str, messages: list[dict], step_name: str, **kwargs) -> str:
469+
"""``_llm_call_async`` for a step that writes a wiki page from the response.
470+
471+
Hard-codes ``raise_on_truncation=True`` so a truncated response skips the
472+
write instead of silently persisting a partial page (#148). Use this for
473+
every page-generating call so the guarantee can't be forgotten at a new
474+
call site.
475+
"""
476+
return await _llm_call_async(model, messages, step_name, raise_on_truncation=True, **kwargs)
477+
478+
451479
async def _close_async_llm_clients() -> None:
452480
"""Close LiteLLM's cached async (aiohttp) clients for the current loop.
453481
@@ -465,22 +493,26 @@ async def _close_async_llm_clients() -> None:
465493
logger.debug("litellm async client cleanup failed", exc_info=True)
466494

467495

468-
def _warn_if_truncated(response, step_name: str, max_tokens: int | None) -> None:
469-
"""Emit a warning when the LLM hit the max_tokens cap.
496+
def _warn_if_truncated(response, step_name: str, max_tokens: int | None) -> bool:
497+
"""Warn when the LLM hit the max_tokens cap; return True if it did.
470498
471499
``json_repair`` will silently salvage the truncated prefix, so without
472-
this the caller can't tell a short response from a cut-off one.
500+
this the caller can't tell a short response from a cut-off one. Callers
501+
that write a page from the response can pass ``raise_on_truncation=True``
502+
to ``_llm_call``/`_llm_call_async`` to turn a truncated response into a
503+
skip instead of persisting partial content.
473504
"""
474505
try:
475506
finish_reason = response.choices[0].finish_reason
476507
except (AttributeError, IndexError):
477-
return
508+
return False
478509
if finish_reason != "length":
479-
return
510+
return False
480511
cap = f" (max_tokens={max_tokens})" if max_tokens else ""
481512
logger.warning("LLM [%s] hit length limit%s — output may be truncated.", step_name, cap)
482513
sys.stdout.write(f" [WARN] {step_name} hit length limit{cap} — output may be truncated.\n")
483514
sys.stdout.flush()
515+
return True
484516

485517

486518
def _parse_json(text: str) -> list | dict:
@@ -499,6 +531,46 @@ def _parse_json(text: str) -> list | dict:
499531
return result
500532

501533

534+
def _parse_page_json(text: str) -> dict | None:
535+
"""Parse an LLM page response into a single JSON object.
536+
537+
Unwraps a single-element ``[{...}]`` array (some models wrap the object in
538+
a list). Returns ``None`` when the response is valid JSON of the wrong
539+
shape (empty/multi-element array, list of scalars) so callers skip the page
540+
rather than persisting the raw JSON text as its body. Propagates the
541+
json/ValueError from ``_parse_json`` when the text isn't JSON at all, which
542+
callers catch to fall back to treating ``raw`` as a prose-markdown body.
543+
"""
544+
parsed = _parse_json(text)
545+
if isinstance(parsed, list) and len(parsed) == 1 and isinstance(parsed[0], dict):
546+
parsed = parsed[0]
547+
return parsed if isinstance(parsed, dict) else None
548+
549+
550+
def _page_fields(raw: str) -> tuple[str, str, dict | None]:
551+
"""Map a page LLM response to ``(brief, content, obj)``.
552+
553+
- JSON object (or a single-element ``[{...}]`` array): brief/content come
554+
from it and ``obj`` is the dict (entity callers read ``type`` from it).
555+
- Valid JSON of the wrong shape (multi/empty array, scalar): ``("", "",
556+
None)`` — the empty content makes ``_require_nonempty_content`` skip the
557+
page rather than persisting the raw JSON text as its body.
558+
- Not JSON at all: ``("", raw, None)`` — ``raw`` is written as a
559+
prose-markdown body (the legitimate fallback for models that emit
560+
markdown instead of JSON).
561+
562+
Shared by all four page-generation closures so a new edge case is handled
563+
in one place instead of four near-identical blocks.
564+
"""
565+
try:
566+
obj = _parse_page_json(raw)
567+
except (json.JSONDecodeError, ValueError):
568+
return "", raw, None
569+
if obj is None:
570+
return "", "", None
571+
return obj.get("description", ""), (obj.get("content") or ""), obj
572+
573+
502574
def _filter_concept_items(items: list, label: str) -> list[dict]:
503575
"""Keep only dicts that carry a non-empty ``name``; warn about anything else."""
504576
if not isinstance(items, list):
@@ -1740,7 +1812,7 @@ async def _gen_create(concept: dict) -> tuple[str, str, bool, str]:
17401812
name = concept["name"]
17411813
title = concept.get("title", name)
17421814
async with semaphore:
1743-
raw = await _llm_call_async(
1815+
raw = await _llm_call_page_async(
17441816
model,
17451817
[
17461818
system_msg,
@@ -1759,17 +1831,7 @@ async def _gen_create(concept: dict) -> tuple[str, str, bool, str]:
17591831
f"concept: {name}",
17601832
response_format=_JSON_RESPONSE_FORMAT,
17611833
)
1762-
try:
1763-
parsed = _parse_json(raw)
1764-
brief = parsed.get("description", "")
1765-
# Parse succeeded: do NOT fall back to ``raw`` (the JSON string).
1766-
# An empty/None ``content`` field yields "" so
1767-
# ``_require_nonempty_content`` raises and the page is skipped,
1768-
# rather than writing the raw JSON as the markdown body.
1769-
content = parsed.get("content") or ""
1770-
except (json.JSONDecodeError, ValueError):
1771-
# Parse FAILED: ``raw`` is the legitimate non-JSON body fallback.
1772-
brief, content = "", raw
1834+
brief, content, _ = _page_fields(raw)
17731835
_require_nonempty_content(content, name)
17741836
return name, content, False, brief
17751837

@@ -1784,7 +1846,7 @@ async def _gen_update(concept: dict) -> tuple[str, str, bool, str]:
17841846
else:
17851847
existing_content = "(page not found — create from scratch)"
17861848
async with semaphore:
1787-
raw = await _llm_call_async(
1849+
raw = await _llm_call_page_async(
17881850
model,
17891851
[
17901852
system_msg,
@@ -1803,14 +1865,7 @@ async def _gen_update(concept: dict) -> tuple[str, str, bool, str]:
18031865
f"update: {name}",
18041866
response_format=_JSON_RESPONSE_FORMAT,
18051867
)
1806-
try:
1807-
parsed = _parse_json(raw)
1808-
brief = parsed.get("description", "")
1809-
# Parse succeeded: do NOT fall back to ``raw`` (the JSON string).
1810-
content = parsed.get("content") or ""
1811-
except (json.JSONDecodeError, ValueError):
1812-
# Parse FAILED: ``raw`` is the legitimate non-JSON body fallback.
1813-
brief, content = "", raw
1868+
brief, content, _ = _page_fields(raw)
18141869
_require_nonempty_content(content, name)
18151870
return name, content, True, brief
18161871

@@ -1819,7 +1874,7 @@ async def _gen_entity_create(ent: dict) -> tuple[str, str, str, str]:
18191874
title = ent.get("title", name)
18201875
etype = ent.get("type", "other")
18211876
async with semaphore:
1822-
raw = await _llm_call_async(
1877+
raw = await _llm_call_page_async(
18231878
model,
18241879
[
18251880
system_msg,
@@ -1838,15 +1893,8 @@ async def _gen_entity_create(ent: dict) -> tuple[str, str, str, str]:
18381893
f"entity: {name}",
18391894
response_format=_JSON_RESPONSE_FORMAT,
18401895
)
1841-
try:
1842-
parsed = _parse_json(raw)
1843-
brief = parsed.get("description", "")
1844-
etype_out = parsed.get("type") if parsed.get("type") in valid_types else etype
1845-
# Parse succeeded: do NOT fall back to ``raw`` (the JSON string).
1846-
content = parsed.get("content") or ""
1847-
except (json.JSONDecodeError, ValueError):
1848-
# Parse FAILED: ``raw`` is the legitimate non-JSON body fallback.
1849-
brief, etype_out, content = "", etype, raw
1896+
brief, content, obj = _page_fields(raw)
1897+
etype_out = obj.get("type") if obj and obj.get("type") in valid_types else etype
18501898
_require_nonempty_content(content, name)
18511899
return name, content, brief, etype_out
18521900

@@ -1862,7 +1910,7 @@ async def _gen_entity_update(ent: dict) -> tuple[str, str, str, str]:
18621910
else:
18631911
existing_content = "(page not found — create from scratch)"
18641912
async with semaphore:
1865-
raw = await _llm_call_async(
1913+
raw = await _llm_call_page_async(
18661914
model,
18671915
[
18681916
system_msg,
@@ -1882,15 +1930,8 @@ async def _gen_entity_update(ent: dict) -> tuple[str, str, str, str]:
18821930
f"entity-update: {name}",
18831931
response_format=_JSON_RESPONSE_FORMAT,
18841932
)
1885-
try:
1886-
parsed = _parse_json(raw)
1887-
brief = parsed.get("description", "")
1888-
etype_out = parsed.get("type") if parsed.get("type") in valid_types else etype
1889-
# Parse succeeded: do NOT fall back to ``raw`` (the JSON string).
1890-
content = parsed.get("content") or ""
1891-
except (json.JSONDecodeError, ValueError):
1892-
# Parse FAILED: ``raw`` is the legitimate non-JSON body fallback.
1893-
brief, etype_out, content = "", etype, raw
1933+
brief, content, obj = _page_fields(raw)
1934+
etype_out = obj.get("type") if obj and obj.get("type") in valid_types else etype
18941935
_require_nonempty_content(content, name)
18951936
return name, content, brief, etype_out
18961937

0 commit comments

Comments
 (0)