Edit on GitHub

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]