agent_search_gateway.academic.aggregator
Identifier-centric clustering and deterministic merge for academic papers.
1"""Identifier-centric clustering and deterministic merge for academic papers.""" 2 3from __future__ import annotations 4 5import logging 6from collections.abc import Callable, Iterable 7from dataclasses import dataclass 8from datetime import date 9 10from ..errors import InputFailure 11from ..models import PaperIdentifiers, PaperRecord 12from ..observability import log_event 13from ..providers.contracts import PaperSearchHit 14from ..url_normalization import NormalizedURL, normalize_url 15from .normalization import ( 16 bibliographic_fingerprint, 17 bibliographic_fingerprints_match, 18 normalize_arxiv_id, 19 normalize_doi, 20 normalize_source_id, 21 normalize_source_name, 22 normalize_topic, 23) 24 25_LOGGER = logging.getLogger(__name__) 26_LLM_PREFIX = "llm:" 27 28 29@dataclass(frozen=True, slots=True) 30class _NormalizedHit: 31 source: str 32 source_id: str 33 title: str 34 authors: tuple[str, ...] 35 abstract: str 36 doi: str 37 arxiv_id: str 38 published_date: date | None 39 updated_date: date | None 40 url: NormalizedURL 41 pdf_url: NormalizedURL | None 42 venue: str 43 topics: tuple[str, ...] 44 citation_count: int | None 45 is_open_access: bool | None 46 oa_status: str 47 license: str 48 origin_rank: tuple[int, int] 49 precedence: tuple[int, int, str, str] 50 51 @property 52 def source_key(self) -> tuple[str, str]: 53 return self.source, self.source_id 54 55 @property 56 def fingerprint(self) -> tuple[str, frozenset[str], int] | None: 57 return bibliographic_fingerprint(self.title, self.authors, self.published_date) 58 59 60@dataclass(slots=True) 61class _Cluster: 62 candidates: list[_NormalizedHit] 63 64 @property 65 def dois(self) -> set[str]: 66 return {candidate.doi for candidate in self.candidates if candidate.doi} 67 68 @property 69 def arxiv_ids(self) -> set[str]: 70 return {candidate.arxiv_id for candidate in self.candidates if candidate.arxiv_id} 71 72 @property 73 def source_keys(self) -> set[tuple[str, str]]: 74 return {candidate.source_key for candidate in self.candidates} 75 76 @property 77 def origin_rank(self) -> tuple[int, int, str, str]: 78 best = min( 79 self.candidates, 80 key=lambda candidate: ( 81 *candidate.origin_rank, 82 candidate.source, 83 candidate.source_id, 84 ), 85 ) 86 return (*best.origin_rank, best.source, best.source_id) 87 88 @property 89 def has_global_strong_id(self) -> bool: 90 return bool(self.dois or self.arxiv_ids) 91 92 93class PaperAggregator: 94 """Normalize candidates, resolve identity, and merge fields deterministically.""" 95 96 def __init__(self, provider_priority: Iterable[str]) -> None: 97 normalized_priority = tuple(normalize_source_name(source) for source in provider_priority) 98 self._priority = { 99 source: index for index, source in enumerate(normalized_priority) if source 100 } 101 self._fallback_rank = len(self._priority) 102 103 def aggregate(self, hits: Iterable[PaperSearchHit]) -> list[PaperRecord]: 104 normalized = self._normalize_hits(list(hits)) 105 clusters: list[_Cluster] = [] 106 for candidate in sorted(normalized, key=self._processing_key): 107 self._attach_candidate(clusters, candidate) 108 ordered_clusters = sorted(clusters, key=lambda cluster: cluster.origin_rank) 109 return [self._merge_cluster(cluster) for cluster in ordered_clusters] 110 111 def _normalize_hits(self, hits: list[PaperSearchHit]) -> list[_NormalizedHit]: 112 source_positions: dict[str, int] = {} 113 normalized: list[_NormalizedHit] = [] 114 for hit in hits: 115 source = normalize_source_name(hit.source) if isinstance(hit.source, str) else "" 116 discovery_index = source_positions.get(source, 0) 117 source_positions[source] = discovery_index + 1 118 candidate = self._normalize_hit(hit, source, discovery_index) 119 if candidate is not None: 120 normalized.append(candidate) 121 return normalized 122 123 def _normalize_hit( 124 self, 125 hit: PaperSearchHit, 126 source: str, 127 discovery_index: int, 128 ) -> _NormalizedHit | None: 129 if not source: 130 self._reject("unknown", "invalid_record_shape") 131 return None 132 title = normalize_topic(hit.title) if isinstance(hit.title, str) else "" 133 if not title: 134 self._reject(source, "missing_title") 135 return None 136 source_identity = normalize_source_id(source, hit.source_id) 137 if source_identity is None: 138 reason = "missing_source_id" if not str(hit.source_id).strip() else "invalid_identifier" 139 self._reject(source, reason) 140 return None 141 _, source_id = source_identity 142 143 doi_valid, canonical_doi = self._optional_identifier(hit.doi, normalize_doi) 144 arxiv_valid, canonical_arxiv = self._optional_identifier( 145 hit.arxiv_id, 146 normalize_arxiv_id, 147 ) 148 if not doi_valid or not arxiv_valid: 149 self._reject(source, "invalid_identifier") 150 return None 151 if source == "crossref": 152 if canonical_doi and canonical_doi != source_id: 153 self._reject(source, "identifier_conflict") 154 return None 155 canonical_doi = source_id 156 if source == "arxiv": 157 if canonical_arxiv and canonical_arxiv != source_id: 158 self._reject(source, "identifier_conflict") 159 return None 160 canonical_arxiv = source_id 161 162 published_valid, published_date = self._date_or_none(hit.published_date) 163 updated_valid, updated_date = self._date_or_none(hit.updated_date) 164 if not published_valid or not updated_valid: 165 self._reject(source, "invalid_date") 166 return None 167 try: 168 url = normalize_url(hit.url) 169 except InputFailure: 170 self._reject(source, "invalid_landing_url") 171 return None 172 pdf_url = self._optional_url(hit.pdf_url) 173 authors = self._strings(hit.authors) 174 topics = self._topics(hit.topics) 175 rank = self._source_rank(source) 176 return _NormalizedHit( 177 source=source, 178 source_id=source_id, 179 title=title, 180 authors=authors, 181 abstract=hit.abstract.strip() if isinstance(hit.abstract, str) else "", 182 doi=canonical_doi, 183 arxiv_id=canonical_arxiv, 184 published_date=published_date, 185 updated_date=updated_date, 186 url=url, 187 pdf_url=pdf_url, 188 venue=normalize_topic(hit.venue) if isinstance(hit.venue, str) else "", 189 topics=topics, 190 citation_count=self._citation_count(hit.citation_count), 191 is_open_access=hit.is_open_access if isinstance(hit.is_open_access, bool) else None, 192 oa_status=normalize_topic(hit.oa_status) if isinstance(hit.oa_status, str) else "", 193 license=normalize_topic(hit.license) if isinstance(hit.license, str) else "", 194 origin_rank=(rank, discovery_index), 195 precedence=(1 if source.startswith(_LLM_PREFIX) else 0, rank, source, source_id), 196 ) 197 198 @staticmethod 199 def _optional_identifier( 200 value: object, 201 normalizer: Callable[[str], str | None], 202 ) -> tuple[bool, str]: 203 if value is None or value == "": 204 return True, "" 205 if not isinstance(value, str): 206 return False, "" 207 normalized = normalizer(value) 208 return (normalized is not None, normalized or "") 209 210 @staticmethod 211 def _date_or_none(value: object) -> tuple[bool, date | None]: 212 if value is None: 213 return True, None 214 if not isinstance(value, date): 215 return False, None 216 return True, value 217 218 @staticmethod 219 def _optional_url(value: object) -> NormalizedURL | None: 220 if not isinstance(value, str) or not value.strip(): 221 return None 222 try: 223 return normalize_url(value) 224 except InputFailure: 225 return None 226 227 @staticmethod 228 def _strings(values: object) -> tuple[str, ...]: 229 if not isinstance(values, (tuple, list)): 230 return () 231 return tuple( 232 cleaned 233 for value in values 234 if isinstance(value, str) and (cleaned := normalize_topic(value)) 235 ) 236 237 @staticmethod 238 def _topics(values: object) -> tuple[str, ...]: 239 if not isinstance(values, (tuple, list)): 240 return () 241 topics: list[str] = [] 242 seen: set[str] = set() 243 for value in values: 244 if not isinstance(value, str): 245 continue 246 normalized = normalize_topic(value) 247 key = normalized.casefold() 248 if normalized and key not in seen: 249 seen.add(key) 250 topics.append(normalized) 251 return tuple(topics) 252 253 @staticmethod 254 def _citation_count(value: object) -> int | None: 255 if isinstance(value, bool) or not isinstance(value, int) or value < 0: 256 return None 257 return value 258 259 def _source_rank(self, source: str) -> int: 260 if source in self._priority: 261 return self._priority[source] 262 if source.startswith(_LLM_PREFIX): 263 return self._fallback_rank + 2 264 return self._fallback_rank + 1 265 266 @staticmethod 267 def _processing_key(candidate: _NormalizedHit) -> tuple[int, int, str, str]: 268 return (*candidate.origin_rank, candidate.source, candidate.source_id) 269 270 def _attach_candidate(self, clusters: list[_Cluster], candidate: _NormalizedHit) -> None: 271 doi_index, arxiv_index, source_index = self._strong_indexes(clusters) 272 references: set[int] = set() 273 if candidate.doi in doi_index: 274 references.add(doi_index[candidate.doi]) 275 if candidate.arxiv_id in arxiv_index: 276 references.add(arxiv_index[candidate.arxiv_id]) 277 if candidate.source_key in source_index: 278 references.add(source_index[candidate.source_key]) 279 280 if not references and not (candidate.doi or candidate.arxiv_id): 281 weak_index = self._weak_index(clusters) 282 fingerprint = candidate.fingerprint 283 if fingerprint is not None: 284 for index in weak_index.get((fingerprint[0], fingerprint[2]), ()): 285 cluster = clusters[index] 286 if cluster.has_global_strong_id: 287 continue 288 if any( 289 bibliographic_fingerprints_match(fingerprint, existing.fingerprint) 290 for existing in cluster.candidates 291 ): 292 references.add(index) 293 294 referenced = [clusters[index] for index in sorted(references)] 295 if self._strong_conflict(referenced, candidate): 296 self._reject(candidate.source, "identifier_conflict") 297 return 298 if not referenced: 299 clusters.append(_Cluster([candidate])) 300 return 301 302 merged_candidates = [candidate] 303 for cluster in referenced: 304 merged_candidates.extend(cluster.candidates) 305 for index in sorted(references, reverse=True): 306 del clusters[index] 307 clusters.append(_Cluster(merged_candidates)) 308 if len(referenced) > 1: 309 log_event( 310 _LOGGER, 311 logging.DEBUG, 312 "paper_clusters_merged", 313 reason="bridged_identity", 314 clusters=len(referenced), 315 ) 316 317 @staticmethod 318 def _strong_indexes( 319 clusters: list[_Cluster], 320 ) -> tuple[dict[str, int], dict[str, int], dict[tuple[str, str], int]]: 321 doi_index: dict[str, int] = {} 322 arxiv_index: dict[str, int] = {} 323 source_index: dict[tuple[str, str], int] = {} 324 for index, cluster in enumerate(clusters): 325 for doi in cluster.dois: 326 doi_index[doi] = index 327 for arxiv_id in cluster.arxiv_ids: 328 arxiv_index[arxiv_id] = index 329 for source_key in cluster.source_keys: 330 source_index[source_key] = index 331 return doi_index, arxiv_index, source_index 332 333 @staticmethod 334 def _weak_index(clusters: list[_Cluster]) -> dict[tuple[str, int], list[int]]: 335 index: dict[tuple[str, int], list[int]] = {} 336 for cluster_index, cluster in enumerate(clusters): 337 if cluster.has_global_strong_id: 338 continue 339 for candidate in cluster.candidates: 340 fingerprint = candidate.fingerprint 341 if fingerprint is not None: 342 index.setdefault((fingerprint[0], fingerprint[2]), []).append(cluster_index) 343 return index 344 345 @staticmethod 346 def _strong_conflict(referenced: list[_Cluster], candidate: _NormalizedHit) -> bool: 347 dois = {candidate.doi} if candidate.doi else set() 348 arxiv_ids = {candidate.arxiv_id} if candidate.arxiv_id else set() 349 for cluster in referenced: 350 dois.update(cluster.dois) 351 arxiv_ids.update(cluster.arxiv_ids) 352 return len(dois) > 1 or len(arxiv_ids) > 1 353 354 def _merge_cluster(self, cluster: _Cluster) -> PaperRecord: 355 candidates = sorted(cluster.candidates, key=lambda candidate: candidate.precedence) 356 title = candidates[0].title 357 authors = next((candidate.authors for candidate in candidates if candidate.authors), ()) 358 abstract = next( 359 (candidate.abstract for candidate in candidates if candidate.abstract), 360 "", 361 ) 362 published_date = next( 363 ( 364 candidate.published_date 365 for candidate in candidates 366 if candidate.published_date is not None 367 ), 368 None, 369 ) 370 updated_date = next( 371 ( 372 candidate.updated_date 373 for candidate in candidates 374 if candidate.updated_date is not None 375 ), 376 None, 377 ) 378 url = candidates[0].url 379 pdf_url = next( 380 (candidate.pdf_url for candidate in candidates if candidate.pdf_url is not None), 381 None, 382 ) 383 venue = next((candidate.venue for candidate in candidates if candidate.venue), "") 384 is_open_access = next( 385 ( 386 candidate.is_open_access 387 for candidate in candidates 388 if candidate.is_open_access is not None 389 ), 390 None, 391 ) 392 oa_status = next( 393 (candidate.oa_status for candidate in candidates if candidate.oa_status), 394 "", 395 ) 396 license_value = next( 397 (candidate.license for candidate in candidates if candidate.license), 398 "", 399 ) 400 return PaperRecord( 401 title=title, 402 authors=authors, 403 abstract=abstract, 404 identifiers=self._identifiers(candidates), 405 published_date=published_date, 406 updated_date=updated_date, 407 url=url, 408 pdf_url=pdf_url, 409 venue=venue, 410 topics=self._stable_topics(candidates), 411 citation_counts=self._citation_counts(candidates), 412 is_open_access=is_open_access, 413 oa_status=oa_status, 414 license=license_value, 415 sources=self._sources(candidates), 416 ) 417 418 @staticmethod 419 def _identifiers(candidates: list[_NormalizedHit]) -> PaperIdentifiers: 420 doi = next((candidate.doi for candidate in candidates if candidate.doi), "") 421 arxiv_id = next((candidate.arxiv_id for candidate in candidates if candidate.arxiv_id), "") 422 native = { 423 source: next( 424 (candidate.source_id for candidate in candidates if candidate.source == source), 425 "", 426 ) 427 for source in ("semantic_scholar", "openalex", "dblp", "core") 428 } 429 return PaperIdentifiers( 430 doi=doi, 431 arxiv_id=arxiv_id, 432 semantic_scholar_id=native["semantic_scholar"], 433 openalex_id=native["openalex"], 434 dblp_key=native["dblp"], 435 core_id=native["core"], 436 ) 437 438 @staticmethod 439 def _stable_topics(candidates: list[_NormalizedHit]) -> tuple[str, ...]: 440 topics: list[str] = [] 441 seen: set[str] = set() 442 for candidate in candidates: 443 for topic in candidate.topics: 444 key = topic.casefold() 445 if key not in seen: 446 seen.add(key) 447 topics.append(topic) 448 return tuple(topics) 449 450 @staticmethod 451 def _citation_counts(candidates: list[_NormalizedHit]) -> dict[str, int]: 452 counts: dict[str, int] = {} 453 for candidate in candidates: 454 if candidate.citation_count is not None: 455 counts[candidate.source] = max( 456 counts.get(candidate.source, 0), 457 candidate.citation_count, 458 ) 459 return counts 460 461 @staticmethod 462 def _sources(candidates: list[_NormalizedHit]) -> tuple[str, ...]: 463 return tuple(dict.fromkeys(candidate.source for candidate in candidates)) 464 465 @staticmethod 466 def _reject(provider: str, reason: str) -> None: 467 log_event( 468 _LOGGER, 469 logging.DEBUG, 470 "paper_candidate_rejected", 471 provider=provider, 472 reason=reason, 473 )
class
PaperAggregator:
94class PaperAggregator: 95 """Normalize candidates, resolve identity, and merge fields deterministically.""" 96 97 def __init__(self, provider_priority: Iterable[str]) -> None: 98 normalized_priority = tuple(normalize_source_name(source) for source in provider_priority) 99 self._priority = { 100 source: index for index, source in enumerate(normalized_priority) if source 101 } 102 self._fallback_rank = len(self._priority) 103 104 def aggregate(self, hits: Iterable[PaperSearchHit]) -> list[PaperRecord]: 105 normalized = self._normalize_hits(list(hits)) 106 clusters: list[_Cluster] = [] 107 for candidate in sorted(normalized, key=self._processing_key): 108 self._attach_candidate(clusters, candidate) 109 ordered_clusters = sorted(clusters, key=lambda cluster: cluster.origin_rank) 110 return [self._merge_cluster(cluster) for cluster in ordered_clusters] 111 112 def _normalize_hits(self, hits: list[PaperSearchHit]) -> list[_NormalizedHit]: 113 source_positions: dict[str, int] = {} 114 normalized: list[_NormalizedHit] = [] 115 for hit in hits: 116 source = normalize_source_name(hit.source) if isinstance(hit.source, str) else "" 117 discovery_index = source_positions.get(source, 0) 118 source_positions[source] = discovery_index + 1 119 candidate = self._normalize_hit(hit, source, discovery_index) 120 if candidate is not None: 121 normalized.append(candidate) 122 return normalized 123 124 def _normalize_hit( 125 self, 126 hit: PaperSearchHit, 127 source: str, 128 discovery_index: int, 129 ) -> _NormalizedHit | None: 130 if not source: 131 self._reject("unknown", "invalid_record_shape") 132 return None 133 title = normalize_topic(hit.title) if isinstance(hit.title, str) else "" 134 if not title: 135 self._reject(source, "missing_title") 136 return None 137 source_identity = normalize_source_id(source, hit.source_id) 138 if source_identity is None: 139 reason = "missing_source_id" if not str(hit.source_id).strip() else "invalid_identifier" 140 self._reject(source, reason) 141 return None 142 _, source_id = source_identity 143 144 doi_valid, canonical_doi = self._optional_identifier(hit.doi, normalize_doi) 145 arxiv_valid, canonical_arxiv = self._optional_identifier( 146 hit.arxiv_id, 147 normalize_arxiv_id, 148 ) 149 if not doi_valid or not arxiv_valid: 150 self._reject(source, "invalid_identifier") 151 return None 152 if source == "crossref": 153 if canonical_doi and canonical_doi != source_id: 154 self._reject(source, "identifier_conflict") 155 return None 156 canonical_doi = source_id 157 if source == "arxiv": 158 if canonical_arxiv and canonical_arxiv != source_id: 159 self._reject(source, "identifier_conflict") 160 return None 161 canonical_arxiv = source_id 162 163 published_valid, published_date = self._date_or_none(hit.published_date) 164 updated_valid, updated_date = self._date_or_none(hit.updated_date) 165 if not published_valid or not updated_valid: 166 self._reject(source, "invalid_date") 167 return None 168 try: 169 url = normalize_url(hit.url) 170 except InputFailure: 171 self._reject(source, "invalid_landing_url") 172 return None 173 pdf_url = self._optional_url(hit.pdf_url) 174 authors = self._strings(hit.authors) 175 topics = self._topics(hit.topics) 176 rank = self._source_rank(source) 177 return _NormalizedHit( 178 source=source, 179 source_id=source_id, 180 title=title, 181 authors=authors, 182 abstract=hit.abstract.strip() if isinstance(hit.abstract, str) else "", 183 doi=canonical_doi, 184 arxiv_id=canonical_arxiv, 185 published_date=published_date, 186 updated_date=updated_date, 187 url=url, 188 pdf_url=pdf_url, 189 venue=normalize_topic(hit.venue) if isinstance(hit.venue, str) else "", 190 topics=topics, 191 citation_count=self._citation_count(hit.citation_count), 192 is_open_access=hit.is_open_access if isinstance(hit.is_open_access, bool) else None, 193 oa_status=normalize_topic(hit.oa_status) if isinstance(hit.oa_status, str) else "", 194 license=normalize_topic(hit.license) if isinstance(hit.license, str) else "", 195 origin_rank=(rank, discovery_index), 196 precedence=(1 if source.startswith(_LLM_PREFIX) else 0, rank, source, source_id), 197 ) 198 199 @staticmethod 200 def _optional_identifier( 201 value: object, 202 normalizer: Callable[[str], str | None], 203 ) -> tuple[bool, str]: 204 if value is None or value == "": 205 return True, "" 206 if not isinstance(value, str): 207 return False, "" 208 normalized = normalizer(value) 209 return (normalized is not None, normalized or "") 210 211 @staticmethod 212 def _date_or_none(value: object) -> tuple[bool, date | None]: 213 if value is None: 214 return True, None 215 if not isinstance(value, date): 216 return False, None 217 return True, value 218 219 @staticmethod 220 def _optional_url(value: object) -> NormalizedURL | None: 221 if not isinstance(value, str) or not value.strip(): 222 return None 223 try: 224 return normalize_url(value) 225 except InputFailure: 226 return None 227 228 @staticmethod 229 def _strings(values: object) -> tuple[str, ...]: 230 if not isinstance(values, (tuple, list)): 231 return () 232 return tuple( 233 cleaned 234 for value in values 235 if isinstance(value, str) and (cleaned := normalize_topic(value)) 236 ) 237 238 @staticmethod 239 def _topics(values: object) -> tuple[str, ...]: 240 if not isinstance(values, (tuple, list)): 241 return () 242 topics: list[str] = [] 243 seen: set[str] = set() 244 for value in values: 245 if not isinstance(value, str): 246 continue 247 normalized = normalize_topic(value) 248 key = normalized.casefold() 249 if normalized and key not in seen: 250 seen.add(key) 251 topics.append(normalized) 252 return tuple(topics) 253 254 @staticmethod 255 def _citation_count(value: object) -> int | None: 256 if isinstance(value, bool) or not isinstance(value, int) or value < 0: 257 return None 258 return value 259 260 def _source_rank(self, source: str) -> int: 261 if source in self._priority: 262 return self._priority[source] 263 if source.startswith(_LLM_PREFIX): 264 return self._fallback_rank + 2 265 return self._fallback_rank + 1 266 267 @staticmethod 268 def _processing_key(candidate: _NormalizedHit) -> tuple[int, int, str, str]: 269 return (*candidate.origin_rank, candidate.source, candidate.source_id) 270 271 def _attach_candidate(self, clusters: list[_Cluster], candidate: _NormalizedHit) -> None: 272 doi_index, arxiv_index, source_index = self._strong_indexes(clusters) 273 references: set[int] = set() 274 if candidate.doi in doi_index: 275 references.add(doi_index[candidate.doi]) 276 if candidate.arxiv_id in arxiv_index: 277 references.add(arxiv_index[candidate.arxiv_id]) 278 if candidate.source_key in source_index: 279 references.add(source_index[candidate.source_key]) 280 281 if not references and not (candidate.doi or candidate.arxiv_id): 282 weak_index = self._weak_index(clusters) 283 fingerprint = candidate.fingerprint 284 if fingerprint is not None: 285 for index in weak_index.get((fingerprint[0], fingerprint[2]), ()): 286 cluster = clusters[index] 287 if cluster.has_global_strong_id: 288 continue 289 if any( 290 bibliographic_fingerprints_match(fingerprint, existing.fingerprint) 291 for existing in cluster.candidates 292 ): 293 references.add(index) 294 295 referenced = [clusters[index] for index in sorted(references)] 296 if self._strong_conflict(referenced, candidate): 297 self._reject(candidate.source, "identifier_conflict") 298 return 299 if not referenced: 300 clusters.append(_Cluster([candidate])) 301 return 302 303 merged_candidates = [candidate] 304 for cluster in referenced: 305 merged_candidates.extend(cluster.candidates) 306 for index in sorted(references, reverse=True): 307 del clusters[index] 308 clusters.append(_Cluster(merged_candidates)) 309 if len(referenced) > 1: 310 log_event( 311 _LOGGER, 312 logging.DEBUG, 313 "paper_clusters_merged", 314 reason="bridged_identity", 315 clusters=len(referenced), 316 ) 317 318 @staticmethod 319 def _strong_indexes( 320 clusters: list[_Cluster], 321 ) -> tuple[dict[str, int], dict[str, int], dict[tuple[str, str], int]]: 322 doi_index: dict[str, int] = {} 323 arxiv_index: dict[str, int] = {} 324 source_index: dict[tuple[str, str], int] = {} 325 for index, cluster in enumerate(clusters): 326 for doi in cluster.dois: 327 doi_index[doi] = index 328 for arxiv_id in cluster.arxiv_ids: 329 arxiv_index[arxiv_id] = index 330 for source_key in cluster.source_keys: 331 source_index[source_key] = index 332 return doi_index, arxiv_index, source_index 333 334 @staticmethod 335 def _weak_index(clusters: list[_Cluster]) -> dict[tuple[str, int], list[int]]: 336 index: dict[tuple[str, int], list[int]] = {} 337 for cluster_index, cluster in enumerate(clusters): 338 if cluster.has_global_strong_id: 339 continue 340 for candidate in cluster.candidates: 341 fingerprint = candidate.fingerprint 342 if fingerprint is not None: 343 index.setdefault((fingerprint[0], fingerprint[2]), []).append(cluster_index) 344 return index 345 346 @staticmethod 347 def _strong_conflict(referenced: list[_Cluster], candidate: _NormalizedHit) -> bool: 348 dois = {candidate.doi} if candidate.doi else set() 349 arxiv_ids = {candidate.arxiv_id} if candidate.arxiv_id else set() 350 for cluster in referenced: 351 dois.update(cluster.dois) 352 arxiv_ids.update(cluster.arxiv_ids) 353 return len(dois) > 1 or len(arxiv_ids) > 1 354 355 def _merge_cluster(self, cluster: _Cluster) -> PaperRecord: 356 candidates = sorted(cluster.candidates, key=lambda candidate: candidate.precedence) 357 title = candidates[0].title 358 authors = next((candidate.authors for candidate in candidates if candidate.authors), ()) 359 abstract = next( 360 (candidate.abstract for candidate in candidates if candidate.abstract), 361 "", 362 ) 363 published_date = next( 364 ( 365 candidate.published_date 366 for candidate in candidates 367 if candidate.published_date is not None 368 ), 369 None, 370 ) 371 updated_date = next( 372 ( 373 candidate.updated_date 374 for candidate in candidates 375 if candidate.updated_date is not None 376 ), 377 None, 378 ) 379 url = candidates[0].url 380 pdf_url = next( 381 (candidate.pdf_url for candidate in candidates if candidate.pdf_url is not None), 382 None, 383 ) 384 venue = next((candidate.venue for candidate in candidates if candidate.venue), "") 385 is_open_access = next( 386 ( 387 candidate.is_open_access 388 for candidate in candidates 389 if candidate.is_open_access is not None 390 ), 391 None, 392 ) 393 oa_status = next( 394 (candidate.oa_status for candidate in candidates if candidate.oa_status), 395 "", 396 ) 397 license_value = next( 398 (candidate.license for candidate in candidates if candidate.license), 399 "", 400 ) 401 return PaperRecord( 402 title=title, 403 authors=authors, 404 abstract=abstract, 405 identifiers=self._identifiers(candidates), 406 published_date=published_date, 407 updated_date=updated_date, 408 url=url, 409 pdf_url=pdf_url, 410 venue=venue, 411 topics=self._stable_topics(candidates), 412 citation_counts=self._citation_counts(candidates), 413 is_open_access=is_open_access, 414 oa_status=oa_status, 415 license=license_value, 416 sources=self._sources(candidates), 417 ) 418 419 @staticmethod 420 def _identifiers(candidates: list[_NormalizedHit]) -> PaperIdentifiers: 421 doi = next((candidate.doi for candidate in candidates if candidate.doi), "") 422 arxiv_id = next((candidate.arxiv_id for candidate in candidates if candidate.arxiv_id), "") 423 native = { 424 source: next( 425 (candidate.source_id for candidate in candidates if candidate.source == source), 426 "", 427 ) 428 for source in ("semantic_scholar", "openalex", "dblp", "core") 429 } 430 return PaperIdentifiers( 431 doi=doi, 432 arxiv_id=arxiv_id, 433 semantic_scholar_id=native["semantic_scholar"], 434 openalex_id=native["openalex"], 435 dblp_key=native["dblp"], 436 core_id=native["core"], 437 ) 438 439 @staticmethod 440 def _stable_topics(candidates: list[_NormalizedHit]) -> tuple[str, ...]: 441 topics: list[str] = [] 442 seen: set[str] = set() 443 for candidate in candidates: 444 for topic in candidate.topics: 445 key = topic.casefold() 446 if key not in seen: 447 seen.add(key) 448 topics.append(topic) 449 return tuple(topics) 450 451 @staticmethod 452 def _citation_counts(candidates: list[_NormalizedHit]) -> dict[str, int]: 453 counts: dict[str, int] = {} 454 for candidate in candidates: 455 if candidate.citation_count is not None: 456 counts[candidate.source] = max( 457 counts.get(candidate.source, 0), 458 candidate.citation_count, 459 ) 460 return counts 461 462 @staticmethod 463 def _sources(candidates: list[_NormalizedHit]) -> tuple[str, ...]: 464 return tuple(dict.fromkeys(candidate.source for candidate in candidates)) 465 466 @staticmethod 467 def _reject(provider: str, reason: str) -> None: 468 log_event( 469 _LOGGER, 470 logging.DEBUG, 471 "paper_candidate_rejected", 472 provider=provider, 473 reason=reason, 474 )
Normalize candidates, resolve identity, and merge fields deterministically.
PaperAggregator(provider_priority: Iterable[str])
97 def __init__(self, provider_priority: Iterable[str]) -> None: 98 normalized_priority = tuple(normalize_source_name(source) for source in provider_priority) 99 self._priority = { 100 source: index for index, source in enumerate(normalized_priority) if source 101 } 102 self._fallback_rank = len(self._priority)
def
aggregate( self, hits: Iterable[agent_search_gateway.providers.contracts.PaperSearchHit]) -> list[agent_search_gateway.models.PaperRecord]:
104 def aggregate(self, hits: Iterable[PaperSearchHit]) -> list[PaperRecord]: 105 normalized = self._normalize_hits(list(hits)) 106 clusters: list[_Cluster] = [] 107 for candidate in sorted(normalized, key=self._processing_key): 108 self._attach_candidate(clusters, candidate) 109 ordered_clusters = sorted(clusters, key=lambda cluster: cluster.origin_rank) 110 return [self._merge_cluster(cluster) for cluster in ordered_clusters]