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
|
valuation_years: int = 4
|
||||||
code_cited_per_code: int = 3
|
code_cited_per_code: int = 3
|
||||||
code_cited_collections: tuple[str, ...] = ("rules", "comments", "corpus")
|
code_cited_collections: tuple[str, ...] = ("rules", "comments", "corpus")
|
||||||
|
code_cited_max: int = 12
|
||||||
|
|
||||||
|
|
||||||
def parse_hosts(spec: str) -> tuple[tuple[str, float], ...]:
|
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
|
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:
|
def load() -> LlmConfig:
|
||||||
"""Read the [llm] section; env vars override host-ish values."""
|
"""Read the [llm] section; env vars override host-ish values."""
|
||||||
from conf import cfg
|
from conf import cfg
|
||||||
@@ -106,9 +116,10 @@ def load() -> LlmConfig:
|
|||||||
or str(_opt(section, "duckdb_replica", "data/replica/aco.ro.duckdb")),
|
or str(_opt(section, "duckdb_replica", "data/replica/aco.ro.duckdb")),
|
||||||
valuation_years=int(_opt(section, "valuation_years", 4)),
|
valuation_years=int(_opt(section, "valuation_years", 4)),
|
||||||
code_cited_per_code=int(_opt(section, "code_cited_per_code", 3)),
|
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"])
|
_opt(section, "code_cited_collections", ["rules", "comments", "corpus"])
|
||||||
),
|
),
|
||||||
|
code_cited_max=int(_opt(section, "code_cited_max", 12)),
|
||||||
embed_dim=int(section.embed_dim),
|
embed_dim=int(section.embed_dim),
|
||||||
build_ann_index=build_ann_index,
|
build_ann_index=build_ann_index,
|
||||||
pg_host=os.environ.get("LLM_PG_HOST", str(section.pg_host)),
|
pg_host=os.environ.get("LLM_PG_HOST", str(section.pg_host)),
|
||||||
|
|||||||
@@ -240,15 +240,19 @@ def _collect(
|
|||||||
per: int,
|
per: int,
|
||||||
collections: Sequence[str],
|
collections: Sequence[str],
|
||||||
seen: set[tuple[str, str]],
|
seen: set[tuple[str, str]],
|
||||||
) -> list[dict]:
|
) -> dict[str, list[dict]]:
|
||||||
"""Run *sql* once per collection in *collections*, capping hits per
|
"""Run *sql* once per collection in *collections*, capping hits per
|
||||||
*key* value (a code or a family) per collection, deduping across the
|
*key* value (a code or a family) per collection, deduping across the
|
||||||
whole call via the shared *seen* set. *key* names both the query
|
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]
|
upper = [w.upper() for w in wanted]
|
||||||
if not upper:
|
if not upper:
|
||||||
return []
|
return {}
|
||||||
out: list[dict] = []
|
out: dict[str, list[dict]] = {}
|
||||||
for collection in collections:
|
for collection in collections:
|
||||||
try:
|
try:
|
||||||
with engine.begin() as conn:
|
with engine.begin() as conn:
|
||||||
@@ -260,6 +264,7 @@ def _collect(
|
|||||||
log.warning("code-cited sources skipped (%s/%s): %s", collection, key, e)
|
log.warning("code-cited sources skipped (%s/%s): %s", collection, key, e)
|
||||||
continue
|
continue
|
||||||
count: dict[str, int] = {w: 0 for w in upper}
|
count: dict[str, int] = {w: 0 for w in upper}
|
||||||
|
collected: list[dict] = []
|
||||||
for document, md in rows:
|
for document, md in rows:
|
||||||
md = {k: str(v) for k, v in (md or {}).items()}
|
md = {k: str(v) for k, v in (md or {}).items()}
|
||||||
dkey = _dedupe_key(md)
|
dkey = _dedupe_key(md)
|
||||||
@@ -271,7 +276,31 @@ def _collect(
|
|||||||
seen.add(dkey)
|
seen.add(dkey)
|
||||||
for w in hits:
|
for w in hits:
|
||||||
count[w] += 1
|
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
|
return out
|
||||||
|
|
||||||
|
|
||||||
@@ -282,39 +311,55 @@ def code_cited_sources(
|
|||||||
per_code: int,
|
per_code: int,
|
||||||
collections: Sequence[str] = ("rules",),
|
collections: Sequence[str] = ("rules",),
|
||||||
families: Sequence[str] = (),
|
families: Sequence[str] = (),
|
||||||
|
max_total: int = 12,
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
"""Chunks across *collections* whose ``codes`` metadata mention any of
|
"""Chunks across *collections* whose ``codes`` metadata mention any of
|
||||||
*codes* — at most *per_code* per code per collection, newest first,
|
*codes* — at most *per_code* per code per collection, newest first,
|
||||||
deduped on (item_key, p_id) for rules or (item_key, seq) otherwise.
|
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
|
When *families* is non-empty, chunks whose ``families`` metadata
|
||||||
mentions one of them are matched too (same per-collection cap) and
|
mentions one of them are matched too (same per-collection cap) —
|
||||||
appended after the code-cited rows — skipped entirely, no query, when
|
skipped entirely, no query, when *families* is empty. Scores are
|
||||||
*families* is empty. Scores are 0.0: these are additive, not ranked.
|
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:
|
if not codes and not families:
|
||||||
return []
|
return []
|
||||||
seen: set[tuple[str, str]] = set()
|
seen: set[tuple[str, str]] = set()
|
||||||
out = _collect(
|
out = _interleave(
|
||||||
engine,
|
_collect(
|
||||||
_CITED_SQL,
|
|
||||||
"codes",
|
|
||||||
codes,
|
|
||||||
per=per_code,
|
|
||||||
collections=collections,
|
|
||||||
seen=seen,
|
|
||||||
)
|
|
||||||
if families:
|
|
||||||
out += _collect(
|
|
||||||
engine,
|
engine,
|
||||||
_FAMILY_CITED_SQL,
|
_CITED_SQL,
|
||||||
"families",
|
"codes",
|
||||||
families,
|
codes,
|
||||||
per=per_code,
|
per=per_code,
|
||||||
collections=collections,
|
collections=collections,
|
||||||
seen=seen,
|
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
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -156,6 +156,7 @@ def stream_answer(
|
|||||||
per_code=cfg.code_cited_per_code,
|
per_code=cfg.code_cited_per_code,
|
||||||
collections=tuple(cfg.code_cited_collections),
|
collections=tuple(cfg.code_cited_collections),
|
||||||
families=evidence.families,
|
families=evidence.families,
|
||||||
|
max_total=cfg.code_cited_max,
|
||||||
)
|
)
|
||||||
sources = merge_sources(sources, cited)
|
sources = merge_sources(sources, cited)
|
||||||
yield evidence.payload()
|
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)
|
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_per_code = 3 # excerpts literally citing each detected code (or family)
|
||||||
code_cited_collections = ["rules", "comments", "corpus"] # searched in this order
|
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
|
[llm.k_per_kind] # over-fetched ×3 per kind, then re-ranked
|
||||||
comment = 8
|
comment = 8
|
||||||
|
|||||||
@@ -95,7 +95,31 @@ class TestNewKnobs:
|
|||||||
assert cfg.valuation_years == 4
|
assert cfg.valuation_years == 4
|
||||||
assert cfg.code_cited_per_code == 3
|
assert cfg.code_cited_per_code == 3
|
||||||
assert cfg.code_cited_collections == ("rules", "comments", "corpus")
|
assert cfg.code_cited_collections == ("rules", "comments", "corpus")
|
||||||
|
assert cfg.code_cited_max == 12
|
||||||
|
|
||||||
def test_duckdb_replica_env_override(self, monkeypatch):
|
def test_duckdb_replica_env_override(self, monkeypatch):
|
||||||
monkeypatch.setenv("LLM_DUCKDB_REPLICA", "/app/data/replica/aco.ro.duckdb")
|
monkeypatch.setenv("LLM_DUCKDB_REPLICA", "/app/data/replica/aco.ro.duckdb")
|
||||||
assert llm_config.load().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(
|
out = code_cited_sources(
|
||||||
engine, ["G0556"], per_code=2, collections=("rules", "comments")
|
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):
|
def test_dedupe_uses_seq_for_non_rule_kinds(self):
|
||||||
engine, _conn = _multi_engine(
|
engine, _conn = _multi_engine(
|
||||||
@@ -568,6 +570,103 @@ class TestFamilyCitedSources:
|
|||||||
assert [s["snippet"] for s in out] == ["shared"]
|
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:
|
class TestMergeSources:
|
||||||
def test_keeps_order_and_dedupes_by_label(self):
|
def test_keeps_order_and_dedupes_by_label(self):
|
||||||
a = [{"label": "X", "score": 0.9}, {"label": "Y", "score": 0.8}]
|
a = [{"label": "X", "score": 0.9}, {"label": "Y", "score": 0.8}]
|
||||||
|
|||||||
@@ -418,6 +418,7 @@ class TestStreamAnswer:
|
|||||||
per_code=CFG.code_cited_per_code,
|
per_code=CFG.code_cited_per_code,
|
||||||
collections=CFG.code_cited_collections,
|
collections=CFG.code_cited_collections,
|
||||||
families=("APCM",),
|
families=("APCM",),
|
||||||
|
max_total=CFG.code_cited_max,
|
||||||
)
|
)
|
||||||
body = client.stream.call_args.kwargs["json"]
|
body = client.stream.call_args.kwargs["json"]
|
||||||
assert "Valuation (authoritative" in body["messages"][1]["content"]
|
assert "Valuation (authoritative" in body["messages"][1]["content"]
|
||||||
|
|||||||
Reference in New Issue
Block a user