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:
@@ -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)),
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",)
|
||||
|
||||
@@ -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}]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user