Merge pull request #423 from hnshah/ren/preserve-requested-quick-sources
This commit is contained in:
@@ -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],
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
Reference in New Issue
Block a user