Merge pull request #366 from davemorin/feat/324-reddit-json-fallback

feat(web): auto-enrich Reddit URLs from web search via JSON API
This commit is contained in:
Trevin Chow
2026-05-16 23:41:57 -07:00
committed by GitHub
2 changed files with 154 additions and 9 deletions
+70 -8
View File
@@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
import sys
import urllib.parse import urllib.parse
from datetime import datetime from datetime import datetime
from urllib.parse import urlparse from urllib.parse import urlparse
@@ -205,29 +206,90 @@ def web_search(
backend = "parallel" backend = "parallel"
else: else:
return [], {} return [], {}
items: list[dict] = []
artifact: dict = {}
if backend == "brave": if backend == "brave":
key = config.get("BRAVE_API_KEY") key = config.get("BRAVE_API_KEY")
if not key: if not key:
raise RuntimeError("BRAVE_API_KEY is required when web_backend='brave'") raise RuntimeError("BRAVE_API_KEY is required when web_backend='brave'")
return brave_search(query, date_range, key) items, artifact = brave_search(query, date_range, key)
if backend == "exa": elif backend == "exa":
key = config.get("EXA_API_KEY") key = config.get("EXA_API_KEY")
if not key: if not key:
raise RuntimeError("EXA_API_KEY is required when web_backend='exa'") raise RuntimeError("EXA_API_KEY is required when web_backend='exa'")
return exa_search(query, date_range, key) items, artifact = exa_search(query, date_range, key)
if backend == "serper": elif backend == "serper":
key = config.get("SERPER_API_KEY") key = config.get("SERPER_API_KEY")
if not key: if not key:
raise RuntimeError("SERPER_API_KEY is required when web_backend='serper'") raise RuntimeError("SERPER_API_KEY is required when web_backend='serper'")
return serper_search(query, date_range, key) items, artifact = serper_search(query, date_range, key)
if backend == "parallel": elif backend == "parallel":
key = config.get("PARALLEL_API_KEY") key = config.get("PARALLEL_API_KEY")
if not key: if not key:
raise RuntimeError("PARALLEL_API_KEY is required when web_backend='parallel'") raise RuntimeError("PARALLEL_API_KEY is required when web_backend='parallel'")
return parallel_search(query, date_range, key) items, artifact = parallel_search(query, date_range, key)
if backend != "none": elif backend != "none":
raise ValueError(f"Unsupported web backend: {backend!r}") raise ValueError(f"Unsupported web backend: {backend!r}")
else:
return [], {} return [], {}
if items and not _reddit_excluded(config):
items = _enrich_reddit_items(items)
return items, artifact
def _reddit_excluded(config: dict) -> bool:
"""Return True when EXCLUDE_SOURCES contains 'reddit'.
Respects the same suppression knob the pipeline uses for source gating,
so a user who set EXCLUDE_SOURCES=reddit doesn't get Reddit content
smuggled back in via web-search URLs.
"""
raw = (config.get("EXCLUDE_SOURCES") or "").split(",")
return any(s.strip().lower() == "reddit" for s in raw)
def _enrich_reddit_items(items: list[dict]) -> list[dict]:
"""Enrich web search results that are Reddit URLs with thread body and comments.
Claude Code's WebFetch blocks reddit.com, so the model can't retrieve
Reddit content from web search results. This fetches it via the public
JSON API (reddit.com/.../.json) which bypasses that restriction.
Callers should gate this with EXCLUDE_SOURCES=reddit handling (see
`_reddit_excluded`) so a user who explicitly excluded Reddit doesn't
get Reddit content via web-search URLs.
"""
from . import reddit_enrich
from .reddit_enrich import RedditRateLimitError
for item in items:
url = item.get("url", "")
if "reddit.com" not in url or "/comments/" not in url:
continue
try:
thread_data = reddit_enrich.fetch_thread_data(url, timeout=8)
if not thread_data:
continue
parsed = reddit_enrich.parse_thread_data(thread_data)
# selftext lives under parsed["submission"], not at the top level
selftext = (parsed.get("submission") or {}).get("selftext", "")
if selftext:
item["snippet"] = selftext[:2000]
comments = parsed.get("comments", [])
top = reddit_enrich.get_top_comments(comments)
if top:
item["top_comments"] = [
{"score": c.get("score", 0), "excerpt": (c.get("body") or "")[:200]}
for c in top[:5]
]
item["enriched_via"] = "reddit_json_api"
except RedditRateLimitError as exc:
# Stop iterating to avoid flooding more 429s
sys.stderr.write(f"[Web] Reddit rate-limited, halting enrichment: {exc}\n")
break
except Exception as exc:
sys.stderr.write(f"[Web] Reddit enrichment failed for {url}: {exc}\n")
return items
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+83
View File
@@ -190,5 +190,88 @@ class WebSearchDispatchTests(unittest.TestCase):
grounding.web_search("test", ("2026-02-25", "2026-03-27"), {}, backend="google") grounding.web_search("test", ("2026-02-25", "2026-03-27"), {}, backend="google")
class RedditEnrichmentGateTests(unittest.TestCase):
"""EXCLUDE_SOURCES=reddit must suppress the web-search Reddit enrichment.
Otherwise a user who explicitly excluded Reddit would still get Reddit
content smuggled back in via web-search URLs that happen to point at
reddit.com threads.
"""
def test_reddit_excluded_via_exclude_sources_skips_enrichment(self):
config = {"BRAVE_API_KEY": "k", "EXCLUDE_SOURCES": "reddit"}
items = [{"url": "https://www.reddit.com/r/python/comments/abc/title/", "snippet": "original"}]
with patch("lib.grounding.brave_search", return_value=(items, {})), \
patch("lib.grounding._enrich_reddit_items") as enrich_mock:
grounding.web_search("test", ("2026-02-25", "2026-03-27"), config, backend="auto")
enrich_mock.assert_not_called()
def test_reddit_excluded_case_insensitive(self):
for value in ("REDDIT", "Reddit", " reddit ", "x,reddit,y"):
config = {"BRAVE_API_KEY": "k", "EXCLUDE_SOURCES": value}
self.assertTrue(
grounding._reddit_excluded(config),
msg=f"_reddit_excluded should be True for EXCLUDE_SOURCES={value!r}",
)
def test_reddit_not_excluded_when_other_sources_listed(self):
config = {"EXCLUDE_SOURCES": "tiktok,instagram"}
self.assertFalse(grounding._reddit_excluded(config))
def test_enrichment_runs_when_reddit_not_excluded(self):
config = {"BRAVE_API_KEY": "k"}
items = [{"url": "https://www.reddit.com/r/python/comments/abc/title/", "snippet": "original"}]
with patch("lib.grounding.brave_search", return_value=(items, {})), \
patch("lib.grounding._enrich_reddit_items", return_value=items) as enrich_mock:
grounding.web_search("test", ("2026-02-25", "2026-03-27"), config, backend="auto")
enrich_mock.assert_called_once()
class RedditEnrichItemsTests(unittest.TestCase):
"""Direct tests for `_enrich_reddit_items` covering the selftext key path
and the RedditRateLimitError early-exit behavior.
"""
def test_selftext_under_submission_populates_snippet(self):
from lib import reddit_enrich
item = {
"url": "https://www.reddit.com/r/python/comments/abc/title/",
"snippet": "original",
}
parsed = {
"submission": {"selftext": "thread body content"},
"comments": [],
}
with patch.object(reddit_enrich, "fetch_thread_data", return_value={"raw": True}), \
patch.object(reddit_enrich, "parse_thread_data", return_value=parsed):
result = grounding._enrich_reddit_items([item])
self.assertEqual("thread body content", result[0]["snippet"])
self.assertEqual("reddit_json_api", result[0]["enriched_via"])
def test_rate_limit_error_halts_iteration(self):
from lib import reddit_enrich
item1 = {"url": "https://www.reddit.com/r/python/comments/aaa/x/"}
item2 = {"url": "https://www.reddit.com/r/python/comments/bbb/y/"}
def fake_fetch(url, *args, **kwargs):
raise reddit_enrich.RedditRateLimitError(f"429 for {url}")
captured_stderr: list[str] = []
with patch.object(reddit_enrich, "fetch_thread_data", side_effect=fake_fetch) as fetch_mock, \
patch("lib.grounding.sys.stderr.write", side_effect=lambda s: captured_stderr.append(s)):
grounding._enrich_reddit_items([item1, item2])
# Only the first item should have triggered a fetch attempt
self.assertEqual(1, fetch_mock.call_count)
# A stderr message about the rate-limit halt should have been emitted
self.assertTrue(
any("rate-limited" in msg.lower() or "rate limited" in msg.lower() for msg in captured_stderr),
msg=f"Expected a rate-limit stderr message, got: {captured_stderr!r}",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()