add criteria-first discovery search provider

This commit is contained in:
Marco0300
2026-09-03 20:04:08 +02:00
parent e3257f1ce9
commit 725ef0db9f
8 changed files with 202 additions and 16 deletions
+16 -7
View File
@@ -1,8 +1,9 @@
"""Bounded, public-only prospect discovery using an explicit seed allowlist.
"""Bounded, public-only prospect discovery with optional criteria search.
There is deliberately no general web search or arbitrary URL input here: callers provide
at most a small set of public seed pages. Only links found on those seeds become candidate
sites, and subsequent crawling is same-origin, bounded, and SSRF-checked by the scanner.
at most a small set of public seed pages. When seeds are omitted, a configured search
provider supplies bounded seed URLs. Links found on those seeds become candidate sites,
and subsequent crawling is same-origin, bounded, and SSRF-checked by the scanner.
"""
from __future__ import annotations
@@ -12,6 +13,7 @@ from urllib.parse import urljoin, urldefrag, urlparse
from .contact_extractor import extract_contacts
from .website_scanner import MAX_BYTES, MAX_REDIRECTS, _fetch, validate_url
from .search_provider import search as search_provider
MAX_SEEDS = 5
MAX_PAGES = 20
@@ -61,14 +63,21 @@ def _criteria_match(text, criteria):
return not keywords or all(str(k).strip().lower() in haystack for k in keywords if str(k).strip())
def discover(criteria, seed_urls, *, max_pages=MAX_PAGES, max_candidates=MAX_CANDIDATES):
def discover(criteria, seed_urls=None, *, max_pages=MAX_PAGES, max_candidates=MAX_CANDIDATES):
if not isinstance(criteria, dict) or len(criteria) > 20: raise ValueError("invalid_criteria")
if not isinstance(seed_urls, list) or not 0 < len(seed_urls) <= MAX_SEEDS: raise ValueError("seed_urls_required")
try: max_pages = int(max_pages); max_candidates = int(max_candidates)
except (TypeError, ValueError): raise ValueError("invalid_limits")
if not 1 <= max_pages <= MAX_PAGES or not 1 <= max_candidates <= MAX_CANDIDATES: raise ValueError("invalid_limits")
if seed_urls is None:
seeds = search_provider(criteria, max_candidates)
mechanism = "criteria_search_provider"
else:
if not isinstance(seed_urls, list) or not 0 < len(seed_urls) <= MAX_SEEDS: raise ValueError("seed_urls_required")
seeds = list(seed_urls)
mechanism = "explicit_seed_allowlist"
raw_seeds = list(seeds)
seeds = []
for raw in seed_urls:
for raw in raw_seeds:
try: safe = validate_url(raw)
except ValueError as exc: raise ValueError("unsafe_seed_url") from exc
if safe not in seeds: seeds.append(safe)
@@ -125,5 +134,5 @@ def discover(criteria, seed_urls, *, max_pages=MAX_PAGES, max_candidates=MAX_CAN
evidence.append({"kind": "discovery_page", "url": page["url"], "claim": claim, "provenance": "scoped_discovery"})
deduped = {(x["kind"], x["value"]): x for x in contacts}
name = next((x["headings"][0] for x in pages if x["headings"]), next((x["title"] for x in pages if x["title"]), domain))
results.append({"name": name[:200], "website": root, "website_domain": domain, "description": text[:1000], "contacts": list(deduped.values())[:100], "evidence": evidence, "pages": pages, "pages_crawled": len(pages), "provenance": {"mechanism": "explicit_seed_allowlist", "seed_urls": seeds, "root_url": root}})
results.append({"name": name[:200], "website": root, "website_domain": domain, "description": text[:1000], "contacts": list(deduped.values())[:100], "evidence": evidence, "pages": pages, "pages_crawled": len(pages), "provenance": {"mechanism": mechanism, "seed_urls": seeds, "root_url": root}})
return {"candidates": results, "seeds": seeds, "pages_limit": max_pages, "candidate_limit": max_candidates}
+18 -8
View File
@@ -15,6 +15,7 @@ if __package__ in (None, ""):
from app.scoring import DEFAULT_RULES, signals_for_business, evaluate_score, SCORE_VERSION
from app.ai_assistance import generate as generate_ai, input_fingerprint, provider_status, MAX_INPUT_ITEMS, MAX_OUTPUT_CHARS
from app.discovery import discover as scoped_discover
from app.search_provider import provider_status as search_provider_status
from app.config import load_config
else:
from .domain import deduplication_key, deduplicate_businesses, is_suppressed, normalize_business, score_business, normalize_domain, normalize_phone, match_businesses
@@ -25,6 +26,7 @@ else:
from .scoring import DEFAULT_RULES, signals_for_business, evaluate_score, SCORE_VERSION
from .ai_assistance import generate as generate_ai, input_fingerprint, provider_status, MAX_INPUT_ITEMS, MAX_OUTPUT_CHARS
from .discovery import discover as scoped_discover
from .search_provider import provider_status as search_provider_status
from .config import load_config
ORGANIZATION_ID = "demo-tenant"
SCHEMA = Path(__file__).resolve().parents[1] / "schema.sql"
@@ -595,6 +597,7 @@ class ApiHandler(BaseHTTPRequestHandler):
if path=="/api/v1/sources": return self.list_sources(db,org)
if path=="/api/v1/discovery-queries": return self.list_queries(db,org)
if path=="/api/v1/discovery-runs": return self.list_discovery_runs(db,org,parse_qs(parsed.query))
if path=="/api/v1/discovery/provider-status": return self.send_json(200, search_provider_status())
if path=="/api/v1/source-records": return self.list_source_records(db,org,parse_qs(parsed.query))
if path=="/api/v1/jobs": return self.list_jobs(db,org,parse_qs(parsed.query))
if path=="/api/v1/domain-checks": return self.list_domain_checks(db,org,parse_qs(parsed.query))
@@ -805,26 +808,33 @@ class ApiHandler(BaseHTTPRequestHandler):
def create_scoped_discovery(self, payload, db, user):
criteria = payload.get("criteria", {}); seeds = payload.get("seed_urls")
if not isinstance(criteria, dict) or not isinstance(seeds, list) or not seeds: return self.send_json(400, {"error": "seed_urls_required"})
criteria_only = seeds is None
if not isinstance(criteria, dict): return self.send_json(400, {"error": "invalid_criteria"})
if not criteria_only and (not isinstance(seeds, list) or not seeds): return self.send_json(400, {"error": "seed_urls_required"})
if criteria_only:
status = search_provider_status()
if status["status"] != "ready": return self.send_json(503, {"error": status["status"], "provider": status["provider"]})
try:
if len(seeds) > 5 or len(json.dumps(criteria).encode()) > 8192: raise ValueError("invalid_criteria")
if not criteria_only and len(seeds) > 5: raise ValueError("invalid_criteria")
if len(json.dumps(criteria).encode()) > 8192: raise ValueError("invalid_criteria")
if not isinstance(criteria.get("keywords", criteria.get("keyword", [])), (list, str)): raise ValueError("invalid_criteria")
max_pages = int(payload.get("max_pages", 20)); max_candidates = int(payload.get("max_candidates", 50))
if not 1 <= max_pages <= 20 or not 1 <= max_candidates <= 50: raise ValueError("invalid_limits")
except (ValueError, TypeError) as exc: return self.send_json(400, {"error": str(exc) or "invalid_criteria"})
try:
for url in seeds: validate_url(url)
except (ValueError, TypeError):
return self.send_json(400, {"error": "unsafe_seed_url"})
if not criteria_only:
try:
for url in seeds: validate_url(url)
except (ValueError, TypeError):
return self.send_json(400, {"error": "unsafe_seed_url"})
key = str(payload.get("idempotency_key", "")).strip()
if not key or len(key) > 200: return self.send_json(400, {"error": "invalid_idempotency_key"})
job_payload = {"criteria": criteria, "seed_urls": seeds, "max_pages": max_pages, "max_candidates": max_candidates}
job_payload = {"criteria": criteria, "seed_urls": seeds if not criteria_only else None, "max_pages": max_pages, "max_candidates": max_candidates}
result = self.create_job({"type": "scoped_discovery", "payload": job_payload, "idempotency_key": key, "_accepted": True, "_defer_wakeup": True}, db, user)
# create_job has already committed; read its id from the response is not available,
# so resolve by the tenant-scoped idempotency key.
job = db.execute("SELECT * FROM jobs WHERE organization_id=? AND idempotency_key=?", (user["organization_id"], key)).fetchone()
if not db.execute("SELECT id FROM discovery_runs WHERE organization_id=? AND job_id=?", (user["organization_id"], job["id"])).fetchone():
db.execute("INSERT INTO discovery_runs(organization_id,job_id,criteria_json,seed_urls_json) VALUES(?,?,?,?)", (user["organization_id"], job["id"], json.dumps(criteria, sort_keys=True), json.dumps(seeds)))
db.execute("INSERT INTO discovery_runs(organization_id,job_id,criteria_json,seed_urls_json) VALUES(?,?,?,?)", (user["organization_id"], job["id"], json.dumps(criteria, sort_keys=True), json.dumps(seeds if not criteria_only else [])))
self.audit(db, user, "discovery.created", str(job["id"])); db.commit()
getattr(self.server, "job_wakeup", threading.Event()).set()
return result
+91
View File
@@ -0,0 +1,91 @@
"""Fail-closed, allowlisted HTTP JSON search provider for criteria-first discovery."""
from __future__ import annotations
import json
import os
from urllib.parse import urlparse
from urllib.request import Request, urlopen
from .website_scanner import validate_url
MAX_RESULTS = 50
MAX_RESPONSE_BYTES = 64 * 1024
TIMEOUT_SECONDS = 8
class ProviderConfigError(ValueError):
"""Provider is absent or configured in a way that is unsafe to call."""
def _config() -> tuple[str, set[str], str]:
endpoint = os.environ.get("SEARCH_PROVIDER_URL", "").strip()
allowed = {x.strip().lower().rstrip(".") for x in os.environ.get("SEARCH_PROVIDER_ALLOWED_HOSTS", "").split(",") if x.strip()}
api_key = os.environ.get("SEARCH_PROVIDER_API_KEY", "").strip()
return endpoint, allowed, api_key
def _validated_endpoint() -> tuple[str, str, str]:
endpoint, allowed, api_key = _config()
if not endpoint:
raise ProviderConfigError("not_configured")
parsed = urlparse(endpoint)
host = (parsed.hostname or "").lower().rstrip(".")
if parsed.scheme != "https" or not host or parsed.username or parsed.password or parsed.fragment or host not in allowed:
raise ProviderConfigError("unsafe_provider")
return endpoint, host, api_key
def provider_status() -> dict[str, object]:
endpoint = os.environ.get("SEARCH_PROVIDER_URL", "").strip()
if not endpoint:
return {"provider": "generic_http_json", "status": "not_configured", "configured": False, "network_enabled": False}
try:
_, host, _ = _validated_endpoint()
except ProviderConfigError as exc:
return {"provider": "generic_http_json", "status": "unsafe_configured", "configured": False, "network_enabled": False, "error": str(exc)}
return {"provider": "generic_http_json", "status": "ready", "configured": True, "network_enabled": True, "host": host, "max_results": MAX_RESULTS}
def _result_urls(payload: object, limit: int) -> list[str]:
items = payload.get("results", payload.get("items", [])) if isinstance(payload, dict) else []
if not isinstance(items, list):
raise ProviderConfigError("invalid_provider_response")
urls: list[str] = []
for item in items[:limit]:
raw = item.get("url", item.get("website", item.get("link", ""))) if isinstance(item, dict) else item
if not isinstance(raw, str) or not raw.strip():
continue
if urlparse(raw.strip()).scheme != "https":
continue
try:
safe = validate_url(raw.strip())
except (TypeError, ValueError):
continue
if safe not in urls:
urls.append(safe)
return urls
def search(criteria: dict, limit: int) -> list[str]:
endpoint, _, api_key = _validated_endpoint()
try:
bounded = max(1, min(int(limit), MAX_RESULTS))
except (TypeError, ValueError) as exc:
raise ProviderConfigError("invalid_limits") from exc
body = json.dumps({"criteria": criteria, "limit": bounded}, separators=(",", ":"), ensure_ascii=False).encode()
headers = {"Content-Type": "application/json", "Accept": "application/json"}
if api_key:
headers["Authorization"] = "Bearer " + api_key
request = Request(endpoint, data=body, headers=headers, method="POST")
try:
with urlopen(request, timeout=TIMEOUT_SECONDS) as response:
raw = response.read(MAX_RESPONSE_BYTES + 1)
except Exception as exc:
raise ProviderConfigError("provider_unavailable") from exc
if len(raw) > MAX_RESPONSE_BYTES:
raise ProviderConfigError("provider_response_too_large")
try:
payload = json.loads(raw.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise ProviderConfigError("invalid_provider_response") from exc
return _result_urls(payload, bounded)
+21 -1
View File
@@ -14,6 +14,8 @@ from app.main import create_server, hash_password
class ScopedDiscoveryApiTests(unittest.TestCase):
def setUp(self):
self.tmp = TemporaryDirectory()
for key in ('SEARCH_PROVIDER_URL', 'SEARCH_PROVIDER_ALLOWED_HOSTS', 'SEARCH_PROVIDER_API_KEY'):
os.environ.pop(key, None)
os.environ['BOOTSTRAP_ADMIN_EMAIL'] = 'discover-owner@example.test'
os.environ['BOOTSTRAP_ADMIN_PASSWORD'] = 'password'
self.server = create_server('127.0.0.1', 0, self.tmp.name + '/db.sqlite')
@@ -60,10 +62,28 @@ class ScopedDiscoveryApiTests(unittest.TestCase):
self.assertFalse(detail.get('outreach_enabled', False))
def test_requires_bounded_seed_allowlist_and_rejects_ssrf(self):
self.assertEqual(self.request('POST', '/api/v1/discovery', {'criteria': {'keywords': ['x']}})[0], 400)
status, body = self.request('POST', '/api/v1/discovery', {'criteria': {'keywords': ['x']}})
self.assertEqual(status, 503); self.assertEqual(body['error'], 'not_configured')
status, body = self.request('POST', '/api/v1/discovery', {'criteria': {}, 'seed_urls': ['http://127.0.0.1/'], 'idempotency_key': 'bad'})
self.assertEqual(status, 400); self.assertEqual(body['error'], 'unsafe_seed_url')
def test_criteria_first_search_results_flow_through_existing_job_persistence(self):
os.environ['SEARCH_PROVIDER_URL'] = 'https://search.example.test/query'
os.environ['SEARCH_PROVIDER_ALLOWED_HOSTS'] = 'search.example.test'
pages = {'https://acme.test/': {'status': 200, 'final_url': 'https://acme.test/', 'content_type': 'text/html', 'body': b'<title>Acme Solar</title><h1>Acme Solar</h1><p>solar</p>'}}
def fetch(url, **_):
value = pages[url]; return dict(value, redirect_chain=[], elapsed_ms=1, tls=True, certificate_status='valid')
with patch('app.discovery.search_provider', return_value=['https://acme.test/']) as provider, patch('app.discovery._fetch', side_effect=fetch), patch('app.discovery.validate_url', side_effect=lambda url, **_: url), patch('app.main.validate_url', side_effect=lambda url, **_: url):
status, created = self.request('POST', '/api/v1/discovery', {'criteria': {'keywords': ['solar']}, 'max_candidates': 1, 'idempotency_key': 'criteria-1'})
self.assertEqual(status, 202)
for _ in range(50):
_, job = self.request('GET', '/api/v1/jobs/' + str(created['id']))
if job['status'] in ('succeeded', 'failed'): break
time.sleep(.02)
self.assertEqual(job['status'], 'succeeded'); provider.assert_called_once_with({'keywords': ['solar']}, 1)
run = self.request('GET', '/api/v1/discovery-runs')[1]['items'][0]
self.assertEqual(run['seed_urls'], []); self.assertEqual(run['result']['candidates'][0]['provenance']['mechanism'], 'criteria_search_provider')
def test_results_are_tenant_isolated(self):
ph, salt = hash_password('other-password')
db = sqlite3.connect(self.server.db_path); db.execute("INSERT INTO organizations VALUES ('other-tenant','Other',CURRENT_TIMESTAMP)"); db.execute("INSERT INTO users (organization_id,email,password_hash,password_salt,role) VALUES (?,?,?,?,?)", ('other-tenant','other@example.test',ph,salt,'owner')); db.commit(); db.close()
+44
View File
@@ -0,0 +1,44 @@
import os
import unittest
from unittest.mock import patch
from app.search_provider import ProviderConfigError, provider_status, search
class SearchProviderTests(unittest.TestCase):
def tearDown(self):
for key in ("SEARCH_PROVIDER_URL", "SEARCH_PROVIDER_ALLOWED_HOSTS", "SEARCH_PROVIDER_API_KEY"):
os.environ.pop(key, None)
def test_not_configured_is_fail_closed(self):
self.assertEqual(provider_status()["status"], "not_configured")
with self.assertRaisesRegex(ProviderConfigError, "not_configured"):
search({"keywords": ["solar"]}, 5)
def test_unsafe_provider_is_rejected(self):
os.environ["SEARCH_PROVIDER_URL"] = "http://search.example.test/query"
os.environ["SEARCH_PROVIDER_ALLOWED_HOSTS"] = "search.example.test"
self.assertEqual(provider_status()["status"], "unsafe_configured")
with self.assertRaisesRegex(ProviderConfigError, "unsafe_provider"):
search({}, 5)
def test_successful_mocked_search_returns_bounded_https_urls(self):
os.environ["SEARCH_PROVIDER_URL"] = "https://search.example.test/query"
os.environ["SEARCH_PROVIDER_ALLOWED_HOSTS"] = "search.example.test"
response = type("Response", (), {"__enter__": lambda self: self, "__exit__": lambda self, *args: None, "read": lambda self, *_: b'{"results":[{"url":"https://acme.test"},{"url":"http://bad.test"},{"url":"https://acme.test"}]}'})()
with patch("app.search_provider.urlopen", return_value=response), patch("app.search_provider.validate_url", side_effect=lambda url, **_: url):
self.assertEqual(search({"keywords": ["solar"]}, 5), ["https://acme.test"])
def test_limit_is_bounded_and_sent_to_provider(self):
os.environ["SEARCH_PROVIDER_URL"] = "https://search.example.test/query"
os.environ["SEARCH_PROVIDER_ALLOWED_HOSTS"] = "search.example.test"
response = type("Response", (), {"__enter__": lambda self: self, "__exit__": lambda self, *args: None, "read": lambda self, *_: b'{"results":[]}'})()
with patch("app.search_provider.urlopen", return_value=response) as opened:
self.assertEqual(search({}, 500), [])
self.assertEqual(opened.call_args.kwargs["timeout"], 8)
request = opened.call_args.args[0]
self.assertIn(b'"limit":50', request.data)
if __name__ == "__main__":
unittest.main()