Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 14 additions & 1 deletion openfoia/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -4538,6 +4538,10 @@ def _progress(event: str, msg: str) -> None:
rprint(f" Sources used: {', '.join(report.sources_used)}")
rprint(f" Total hits: {report.total_hits}")
rprint(f" Entities flagged: {report.total_flagged}")
source_errors = getattr(report, "source_errors", {})
if source_errors:
details = ", ".join(f"{source} ({error})" for source, error in source_errors.items())
rprint(f" [red]Sources errored: {len(source_errors)} — {details}[/red]")
rprint("=" * 60)

for result in report.results:
Expand Down Expand Up @@ -4567,7 +4571,13 @@ def _progress(event: str, msg: str) -> None:
rprint(f" {hit.url}")

if not report.total_flagged:
rprint("\n [green]No cross-reference hits found.[/green]")
if source_errors:
rprint(
"\n [yellow]No cross-reference hits found among completed source checks; "
"report is incomplete.[/yellow]"
)
else:
rprint("\n [green]No cross-reference hits found.[/green]")

# Export as FollowTheMoney
if ftm:
Expand All @@ -4583,10 +4593,12 @@ def _progress(event: str, msg: str) -> None:
"total_hits": report.total_hits,
"total_flagged": report.total_flagged,
"sources": report.sources_used,
"source_errors": source_errors,
"results": [
{
"entity": r.entity_name,
"type": r.entity_type,
"source_statuses": getattr(r, "source_statuses", {}),
"hits": [
{
"source": h.source,
Expand All @@ -4599,6 +4611,7 @@ def _progress(event: str, msg: str) -> None:
}
for r in report.results
if r.hits
or any(status != "checked" for status in getattr(r, "source_statuses", {}).values())
],
}
output.write_text(json.dumps(report_data, indent=2))
Expand Down
136 changes: 105 additions & 31 deletions openfoia/crossref.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ class CrossRefResult:
entity_type: str
hits: list[CrossRefHit]
sources_checked: list[str]
source_statuses: dict[str, str] = field(default_factory=dict)

@property
def flagged(self) -> bool:
Expand All @@ -59,6 +60,7 @@ class CrossRefReport:
total_hits: int
total_flagged: int
sources_used: list[str]
source_errors: dict[str, str] = field(default_factory=dict)


# Entity types worth cross-referencing (skip dates, money, etc.)
Expand Down Expand Up @@ -92,11 +94,21 @@ class _RateLimited(BaseException):
"""


class _SourceCheckError(RuntimeError):
"""A checker failure that must be visible in the cross-reference report."""

def __init__(self, error: BaseException):
self.error_type = type(error).__name__
super().__init__(self.error_type)


def _check_rate_limit(result: Any) -> None:
"""Raise _RateLimited if the search result indicates a rate limit error."""
err = getattr(result, "error", None)
if err and ("rate limit" in err.lower() or "429" in err.lower()):
raise _RateLimited(err)
if err:
if "rate limit" in err.lower() or "429" in err.lower():
raise _RateLimited(err)
raise _SourceCheckError(RuntimeError(err))


def _deduplicate_entities(entities: list[Any]) -> list[Any]:
Expand Down Expand Up @@ -194,12 +206,14 @@ async def crossref_entities(
)

results: list[CrossRefResult] = []
source_errors: dict[str, str] = {}
# Track sources that hit rate limits — skip them for remaining entities
exhausted_sources: set[str] = set()

for idx, entity in enumerate(targets):
hits: list[CrossRefHit] = []
sources_checked: list[str] = []
source_statuses: dict[str, str] = {}

if on_progress:
on_progress(
Expand All @@ -209,24 +223,38 @@ async def crossref_entities(

for source_name, checker in available_sources.items():
if source_name in exhausted_sources:
source_statuses[source_name] = "skipped(rate-limited)"
continue
sources_checked.append(source_name)
try:
source_hits = await checker(entity.normalized_text, entity.entity_type)
hits.extend(source_hits)
source_statuses[source_name] = "matched" if source_hits else "checked"
except _RateLimited:
logger.warning(
"CrossRef %s rate limited — skipping for remaining entities",
source_name,
)
exhausted_sources.add(source_name)
source_statuses[source_name] = "ERRORED(RateLimited)"
source_errors.setdefault(source_name, "RateLimited")
except _SourceCheckError as exc:
logger.warning(
"CrossRef %s failed: %s",
source_name,
exc.error_type,
)
source_statuses[source_name] = f"ERRORED({exc.error_type})"
source_errors.setdefault(source_name, exc.error_type)
except Exception as e:
logger.warning(
"CrossRef %s failed for '%s': %s",
"CrossRef %s failed: %s",
source_name,
entity.normalized_text,
e,
type(e).__name__,
)
error_type = type(e).__name__
source_statuses[source_name] = f"ERRORED({error_type})"
source_errors.setdefault(source_name, error_type)

# Rate limit: respect each API's documented limits. Sleep the base
# delay times a randomized ~[1.0, 1.5) jitter factor rather than
Expand All @@ -246,6 +274,7 @@ async def crossref_entities(
entity_type=entity.entity_type.value,
hits=hits,
sources_checked=sources_checked,
source_statuses=source_statuses,
)
)

Expand All @@ -258,6 +287,7 @@ async def crossref_entities(
total_hits=total_hits,
total_flagged=len(flagged),
sources_used=list(available_sources.keys()),
source_errors=source_errors,
)


Expand Down Expand Up @@ -327,8 +357,12 @@ async def _check_muckrock(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for req in result.entities:
Expand Down Expand Up @@ -374,8 +408,12 @@ async def _check_opencorporates(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for ent in result.entities:
Expand Down Expand Up @@ -417,8 +455,12 @@ async def _check_sec(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
seen_ciks: set[str] = set()
Expand Down Expand Up @@ -455,9 +497,11 @@ async def _check_icij(name: str, entity_type: EntityType, data_dir: str) -> list
data_path = Path(data_dir)
name_lower = name.lower()

# Search across all ICIJ CSV files
for csv_file in data_path.glob("*.csv"):
try:
# Search across all ICIJ CSV files. A local data read failure makes this
# source incomplete, so surface it to the report rather than treating it
# as an uneventful no-match.
try:
for csv_file in data_path.glob("*.csv"):
with open(csv_file, encoding="utf-8", errors="ignore") as f:
reader = csv.DictReader(f)
for row in reader:
Expand All @@ -480,8 +524,8 @@ async def _check_icij(name: str, entity_type: EntityType, data_dir: str) -> list
)
)
break # one hit per row is enough
except Exception as e:
logger.warning("Failed to search ICIJ file %s: %s", csv_file, e)
except Exception as exc:
raise _SourceCheckError(exc) from exc

return hits[:10] # cap to avoid flooding

Expand All @@ -496,8 +540,12 @@ async def _check_fec(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for contrib in result.entities:
Expand Down Expand Up @@ -529,8 +577,12 @@ async def _check_regulations(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for doc in result.entities:
Expand Down Expand Up @@ -564,8 +616,12 @@ async def _check_govinfo(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for doc in result.entities:
Expand Down Expand Up @@ -604,8 +660,12 @@ async def _check_nonprofits(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for org in result.entities:
Expand Down Expand Up @@ -646,8 +706,12 @@ async def _check_usaspending(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for award in result.entities:
Expand Down Expand Up @@ -687,8 +751,12 @@ async def _check_documentcloud(
try:
result = await adapter.search(name, page_size=5)
_check_rate_limit(result)
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

hits = []
for doc in result.entities:
Expand Down Expand Up @@ -736,11 +804,17 @@ async def _check_opensanctions(
params={"q": name, "limit": 5},
headers={"Accept": "application/json"},
)
if resp.status_code == 429:
raise _RateLimited("OpenSanctions rate limit")
if resp.status_code != 200:
return []
raise _SourceCheckError(RuntimeError(f"HTTP {resp.status_code}"))
data = resp.json()
except Exception:
return []
except _RateLimited:
raise
except _SourceCheckError:
raise
except Exception as exc:
raise _SourceCheckError(exc) from exc

for result in data.get("results", []):
score = result.get("score", 0)
Expand Down
Loading
Loading