feat(llm): cap cited sources with round-robin interleave (max_total)

code_cited_sources gains max_total (default 12, <=0 = unlimited): after
the existing per-collection caps and dedupe, code-cited rows are
round-robin interleaved across the configured collections in order,
then family-cited rows the same way, appended after — so one chatty
collection can no longer crowd the others out ahead of a truncation
cut. New [llm] knob code_cited_max wires it from config into
stream_answer. Also guards code_cited_collections against a bare
scalar string in stack.toml collapsing into a tuple of characters.
This commit is contained in:
kert
2026-09-10 00:55:58 -04:00
parent 3d8a0725d8
commit 31d9f27d11
7 changed files with 207 additions and 25 deletions

View File

@@ -48,6 +48,7 @@ class LlmConfig:
valuation_years: int = 4
code_cited_per_code: int = 3
code_cited_collections: tuple[str, ...] = ("rules", "comments", "corpus")
code_cited_max: int = 12
def parse_hosts(spec: str) -> tuple[tuple[str, float], ...]:
@@ -72,6 +73,15 @@ def _opt(section: Any, key: str, default: Any) -> Any:
return section[key] if key in section else default
def _str_tuple(val: Any) -> tuple[str, ...]:
"""*val* as a tuple of strings — a bare scalar string in stack.toml
(``code_cited_collections = "rules"``) becomes a one-element tuple
rather than ``tuple("rules")`` exploding it into its characters."""
if isinstance(val, str):
return (val,)
return tuple(val)
def load() -> LlmConfig:
"""Read the [llm] section; env vars override host-ish values."""
from conf import cfg
@@ -106,9 +116,10 @@ def load() -> LlmConfig:
or str(_opt(section, "duckdb_replica", "data/replica/aco.ro.duckdb")),
valuation_years=int(_opt(section, "valuation_years", 4)),
code_cited_per_code=int(_opt(section, "code_cited_per_code", 3)),
code_cited_collections=tuple(
code_cited_collections=_str_tuple(
_opt(section, "code_cited_collections", ["rules", "comments", "corpus"])
),
code_cited_max=int(_opt(section, "code_cited_max", 12)),
embed_dim=int(section.embed_dim),
build_ann_index=build_ann_index,
pg_host=os.environ.get("LLM_PG_HOST", str(section.pg_host)),

View File

@@ -240,15 +240,19 @@ def _collect(
per: int,
collections: Sequence[str],
seen: set[tuple[str, str]],
) -> list[dict]:
) -> dict[str, list[dict]]:
"""Run *sql* once per collection in *collections*, capping hits per
*key* value (a code or a family) per collection, deduping across the
whole call via the shared *seen* set. *key* names both the query
parameter (``:codes``/``:families``) and the chunk metadata field."""
parameter (``:codes``/``:families``) and the chunk metadata field.
Returns each collection's own hit list, keyed by collection name (a
collection with no hits is simply absent), so a caller can
round-robin interleave across collections instead of only ever
getting them concatenated."""
upper = [w.upper() for w in wanted]
if not upper:
return []
out: list[dict] = []
return {}
out: dict[str, list[dict]] = {}
for collection in collections:
try:
with engine.begin() as conn:
@@ -260,6 +264,7 @@ def _collect(
log.warning("code-cited sources skipped (%s/%s): %s", collection, key, e)
continue
count: dict[str, int] = {w: 0 for w in upper}
collected: list[dict] = []
for document, md in rows:
md = {k: str(v) for k, v in (md or {}).items()}
dkey = _dedupe_key(md)
@@ -271,7 +276,31 @@ def _collect(
seen.add(dkey)
for w in hits:
count[w] += 1
out.append(as_source(md, document[:250], 0.0))
collected.append(as_source(md, document[:250], 0.0))
if collected:
out[collection] = collected
return out
def _interleave(
per_collection: dict[str, list[dict]], collections: Sequence[str]
) -> list[dict]:
"""Round-robin merge of *per_collection*'s lists in *collections*
order — one row from each collection in turn, looping until every
list is exhausted, so a collection with many hits can't crowd the
others out of the front of the list ahead of a ``max_total`` cut."""
out: list[dict] = []
idx = {c: 0 for c in collections}
progressed = True
while progressed:
progressed = False
for c in collections:
rows = per_collection.get(c, [])
i = idx[c]
if i < len(rows):
out.append(rows[i])
idx[c] = i + 1
progressed = True
return out
@@ -282,39 +311,55 @@ def code_cited_sources(
per_code: int,
collections: Sequence[str] = ("rules",),
families: Sequence[str] = (),
max_total: int = 12,
) -> list[dict]:
"""Chunks across *collections* whose ``codes`` metadata mention any of
*codes* — at most *per_code* per code per collection, newest first,
deduped on (item_key, p_id) for rules or (item_key, seq) otherwise.
Concatenated in the order *collections* is given.
When *families* is non-empty, chunks whose ``families`` metadata
mentions one of them are matched too (same per-collection cap) and
appended after the code-cited rows — skipped entirely, no query, when
*families* is empty. Scores are 0.0: these are additive, not ranked.
mentions one of them are matched too (same per-collection cap) —
skipped entirely, no query, when *families* is empty. Scores are
0.0: these are additive, not ranked.
After the per-collection caps and dedupe above, the code-cited rows
are round-robin interleaved across *collections* in the order given
(one from the first collection, one from the second, … looping back
until every collection's list is exhausted) so no single collection
crowds the others out; the family-cited rows are interleaved the
same way and appended after. The combined list is then truncated to
*max_total* — ``<= 0`` means unlimited.
"""
if not codes and not families:
return []
seen: set[tuple[str, str]] = set()
out = _collect(
engine,
_CITED_SQL,
"codes",
codes,
per=per_code,
collections=collections,
seen=seen,
)
if families:
out += _collect(
out = _interleave(
_collect(
engine,
_FAMILY_CITED_SQL,
"families",
families,
_CITED_SQL,
"codes",
codes,
per=per_code,
collections=collections,
seen=seen,
),
collections,
)
if families:
out += _interleave(
_collect(
engine,
_FAMILY_CITED_SQL,
"families",
families,
per=per_code,
collections=collections,
seen=seen,
),
collections,
)
if max_total > 0:
out = out[:max_total]
return out

View File

@@ -156,6 +156,7 @@ def stream_answer(
per_code=cfg.code_cited_per_code,
collections=tuple(cfg.code_cited_collections),
families=evidence.families,
max_total=cfg.code_cited_max,
)
sources = merge_sources(sources, cited)
yield evidence.payload()

View File

@@ -128,6 +128,7 @@ duckdb_replica = "data/replica/aco.ro.duckdb" # read-only DuckDB replica the c
valuation_years = 4 # final-rule vintages shown per code (plus the newest NPRM)
code_cited_per_code = 3 # excerpts literally citing each detected code (or family)
code_cited_collections = ["rules", "comments", "corpus"] # searched in this order
code_cited_max = 12 # cited excerpts kept after round-robin interleave across collections; <= 0 = unlimited
[llm.k_per_kind] # over-fetched ×3 per kind, then re-ranked
comment = 8

View File

@@ -95,7 +95,31 @@ class TestNewKnobs:
assert cfg.valuation_years == 4
assert cfg.code_cited_per_code == 3
assert cfg.code_cited_collections == ("rules", "comments", "corpus")
assert cfg.code_cited_max == 12
def test_duckdb_replica_env_override(self, monkeypatch):
monkeypatch.setenv("LLM_DUCKDB_REPLICA", "/app/data/replica/aco.ro.duckdb")
assert llm_config.load().duckdb_replica == "/app/data/replica/aco.ro.duckdb"
class TestStrTuple:
"""``code_cited_collections`` must never explode a bare scalar string
into a tuple of its characters (``tuple("rules")``)."""
def test_scalar_string_becomes_one_tuple(self):
assert llm_config._str_tuple("rules") == ("rules",)
def test_list_passes_through_as_tuple(self):
assert llm_config._str_tuple(["rules", "comments"]) == ("rules", "comments")
def test_code_cited_collections_scalar_string_in_toml_is_one_tuple(
self, monkeypatch
):
import conf
from conf import _Cfg
raw = dict(conf.cfg.to_dict())
raw["llm"] = dict(raw["llm"])
raw["llm"]["code_cited_collections"] = "rules"
monkeypatch.setattr(conf, "cfg", _Cfg(raw))
assert llm_config.load().code_cited_collections == ("rules",)

View File

@@ -472,7 +472,9 @@ class TestCodeCitedSourcesMultiCollection:
out = code_cited_sources(
engine, ["G0556"], per_code=2, collections=("rules", "comments")
)
assert [s["snippet"] for s in out] == ["r1", "r2", "c1", "c2"]
# Capped to 2/collection, then round-robin interleaved across
# collections (rules, comments), not simply concatenated.
assert [s["snippet"] for s in out] == ["r1", "c1", "r2", "c2"]
def test_dedupe_uses_seq_for_non_rule_kinds(self):
engine, _conn = _multi_engine(
@@ -568,6 +570,103 @@ class TestFamilyCitedSources:
assert [s["snippet"] for s in out] == ["shared"]
class TestMaxTotal:
"""``max_total`` interleaves code/family rows round-robin across
collections, then truncates (#budget)."""
def test_interleaves_three_collections_round_robin(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556")), ("r2", _rule_md(2, "G0556"))],
"comments": [
("c1", _comment_md(1, "G0556")),
("c2", _comment_md(2, "G0556")),
],
"corpus": [
("x1", _corpus_md(1, codes="G0556", item_key="X1")),
("x2", _corpus_md(2, codes="G0556", item_key="X2")),
],
}
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments", "corpus"),
max_total=0,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "x1", "r2", "c2", "x2"]
def test_family_rows_interleaved_then_appended_after_code_rows(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556"))],
"comments": [("c1", _comment_md(1, "G0556"))],
},
by_family={
"rules": [
(
"f-rules",
_corpus_md(9, codes="", families="CCM", item_key="FR"),
)
],
"comments": [
(
"f-comments",
_corpus_md(10, codes="", families="CCM", item_key="FC"),
)
],
},
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments"),
families=["CCM"],
max_total=0,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "f-rules", "f-comments"]
def test_truncates_to_max_total(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556")), ("r2", _rule_md(2, "G0556"))],
"comments": [
("c1", _comment_md(1, "G0556")),
("c2", _comment_md(2, "G0556")),
],
}
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments"),
max_total=3,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "r2"]
def test_max_total_zero_or_negative_is_unlimited(self):
rows = [(f"r{i}", _rule_md(i, "G0556")) for i in range(1, 15)]
engine, _conn = _multi_engine(by_code={"rules": rows})
for unlimited in (0, -1):
out = code_cited_sources(
engine,
["G0556"],
per_code=20,
collections=("rules",),
max_total=unlimited,
)
assert len(out) == 14
def test_default_max_total_is_twelve(self):
rows = [(f"r{i}", _rule_md(i, "G0556")) for i in range(1, 15)]
engine, _conn = _multi_engine(by_code={"rules": rows})
out = code_cited_sources(engine, ["G0556"], per_code=20, collections=("rules",))
assert len(out) == 12
class TestMergeSources:
def test_keeps_order_and_dedupes_by_label(self):
a = [{"label": "X", "score": 0.9}, {"label": "Y", "score": 0.8}]

View File

@@ -418,6 +418,7 @@ class TestStreamAnswer:
per_code=CFG.code_cited_per_code,
collections=CFG.code_cited_collections,
families=("APCM",),
max_total=CFG.code_cited_max,
)
body = client.stream.call_args.kwargs["json"]
assert "Valuation (authoritative" in body["messages"][1]["content"]