Merge pull request #423 from hnshah/ren/preserve-requested-quick-sources

This commit is contained in:
Trevin Chow
2026-05-22 08:14:46 -07:00
committed by GitHub
2 changed files with 61 additions and 2 deletions
+26 -2
View File
@@ -274,7 +274,15 @@ def _sanitize_plan(
freshness_mode=freshness_mode, freshness_mode=freshness_mode,
cluster_mode=cluster_mode, cluster_mode=cluster_mode,
raw_topic=topic, raw_topic=topic,
subqueries=_normalize_subquery_weights(_trim_subqueries_for_depth(subqueries, intent, depth, eligible_sources)), subqueries=_normalize_subquery_weights(
_trim_subqueries_for_depth(
subqueries,
intent,
depth,
eligible_sources,
requested_sources=requested_sources,
)
),
source_weights=source_weights, source_weights=source_weights,
notes=[str(note).strip() for note in raw.get("notes") or [] if str(note).strip()], notes=[str(note).strip() for note in raw.get("notes") or [] if str(note).strip()],
) )
@@ -307,6 +315,7 @@ def _trim_subqueries_for_depth(
intent: str, intent: str,
depth: str, depth: str,
available_sources: list[str], available_sources: list[str],
requested_sources: list[str] | None = None,
) -> list[schema.SubQuery]: ) -> list[schema.SubQuery]:
# At non-quick depth, expand sources: use capability routing for intents # At non-quick depth, expand sources: use capability routing for intents
# that define it, or all available sources otherwise. The LLM planner may # that define it, or all available sources otherwise. The LLM planner may
@@ -336,6 +345,15 @@ def _trim_subqueries_for_depth(
for subquery in subqueries: for subquery in subqueries:
if depth in {"quick", "default"}: if depth in {"quick", "default"}:
preferred_sources = ranked_sources[:limit] preferred_sources = ranked_sources[:limit]
if requested_sources:
requested = [
source
for source in requested_sources
if source in available_sources and source in subquery.sources
]
for source in requested:
if source not in preferred_sources:
preferred_sources.append(source)
else: else:
preferred_sources = [source for source in ranked_sources if source in subquery.sources][:limit] preferred_sources = [source for source in ranked_sources if source in subquery.sources][:limit]
if len(preferred_sources) < limit: if len(preferred_sources) < limit:
@@ -428,7 +446,13 @@ def _fallback_plan(
cluster_mode=_default_cluster_mode(intent), cluster_mode=_default_cluster_mode(intent),
raw_topic=topic, raw_topic=topic,
subqueries=_normalize_subquery_weights( subqueries=_normalize_subquery_weights(
_trim_subqueries_for_depth(subqueries[:_max_subqueries(intent, topic)], intent, depth, list(source_weights)) _trim_subqueries_for_depth(
subqueries[:_max_subqueries(intent, topic)],
intent,
depth,
list(source_weights),
requested_sources=requested_sources,
)
), ),
source_weights=_normalize_weights(source_weights), source_weights=_normalize_weights(source_weights),
notes=[note], notes=[note],
+35
View File
@@ -92,6 +92,41 @@ class PlannerV3Tests(unittest.TestCase):
self.assertEqual(1, len(plan.subqueries)) self.assertEqual(1, len(plan.subqueries))
self.assertEqual(["reddit", "x"], plan.subqueries[0].sources) self.assertEqual(["reddit", "x"], plan.subqueries[0].sources)
def test_quick_mode_preserves_explicit_requested_sources(self):
raw = {
"intent": "product",
"freshness_mode": "balanced_recent",
"cluster_mode": "debate",
"subqueries": [
{
"label": "primary",
"search_query": "AI coding agents",
"ranking_query": "What are people saying about AI coding agents?",
"sources": ["reddit", "youtube", "grounding", "digg"],
"weight": 1.0,
}
],
}
plan = planner._sanitize_plan(
raw,
"AI coding agents",
["reddit", "youtube", "grounding", "digg"],
["reddit", "youtube", "grounding", "digg"],
"quick",
)
self.assertIn("digg", plan.subqueries[0].sources)
def test_quick_mode_preserves_explicit_requested_sources_in_fallback_plan(self):
plan = planner.plan_query(
topic="AI coding agents",
available_sources=["reddit", "youtube", "github"],
requested_sources=["reddit", "github"],
depth="quick",
provider=None,
model=None,
)
self.assertIn("github", plan.subqueries[0].sources)
def test_default_comparison_uses_all_capable_sources(self): def test_default_comparison_uses_all_capable_sources(self):
plan = planner.plan_query( plan = planner.plan_query(
topic="codex vs claude code", topic="codex vs claude code",