Edit on GitHub

agent_search_gateway.orchestrators.paper

Direct academic paper discovery orchestration and shared paper finalization.

  1"""Direct academic paper discovery orchestration and shared paper finalization."""
  2
  3from __future__ import annotations
  4
  5import asyncio
  6import logging
  7import time
  8from collections.abc import Callable, Sequence
  9
 10from ..academic.aggregator import PaperAggregator
 11from ..academic.enrichment import enrich_paper_records
 12from ..concurrency import ProviderQuotaManager
 13from ..errors import ErrorCode, ExecutionFailure, InputFailure
 14from ..models import PaperRecord
 15from ..observability import elapsed_ms, log_event
 16from ..providers.contracts import AcademicSearchProvider, OAResolver, PaperSearchHit
 17from ..request_ids import validate_request_id
 18from ..result_writer import ResultWriter
 19from ..url_store import URLStore
 20
 21
 22async def finalize_paper_hits(
 23    hits: Sequence[PaperSearchHit],
 24    *,
 25    aggregator: PaperAggregator,
 26    resolver: OAResolver | None,
 27    store: URLStore,
 28) -> list[PaperRecord]:
 29    """Apply the one shared aggregate -> OA enrich -> landing admission policy."""
 30
 31    records = aggregator.aggregate(hits)
 32    enriched = await enrich_paper_records(records, resolver)
 33    for record in enriched:
 34        admission_abstract = record.abstract.strip() or record.title.strip()
 35        store.admit(record.url, admission_abstract)
 36    return enriched
 37
 38
 39class PaperSearchOrchestrator:
 40    def __init__(
 41        self,
 42        *,
 43        providers: Sequence[AcademicSearchProvider],
 44        quotas: ProviderQuotaManager,
 45        aggregator: PaperAggregator,
 46        resolver: OAResolver | None,
 47        store: URLStore,
 48        result_writer: ResultWriter,
 49        logger: logging.Logger | None = None,
 50        monotonic: Callable[[], float] = time.monotonic,
 51    ) -> None:
 52        self.providers = tuple(providers)
 53        self.quotas = quotas
 54        self.aggregator = aggregator
 55        self.resolver = resolver
 56        self.store = store
 57        self.result_writer = result_writer
 58        self._logger = logger or logging.getLogger(__name__)
 59        self._monotonic = monotonic
 60
 61    async def paper_search(self, query: str, *, request_id: str) -> str:
 62        validate_request_id(request_id)
 63        normalized_query = query.strip()
 64        if not normalized_query:
 65            raise InputFailure(ErrorCode.EMPTY_QUERY, "Query must not be empty")
 66        if not self.providers:
 67            raise ExecutionFailure(
 68                ErrorCode.NO_ACADEMIC_SEARCH_PROVIDERS,
 69                "No academic search providers are enabled",
 70            )
 71
 72        outcomes = await asyncio.gather(
 73            *(self._run_provider(provider, normalized_query) for provider in self.providers),
 74            return_exceptions=True,
 75        )
 76        if not any(isinstance(outcome, list) for outcome in outcomes):
 77            raise ExecutionFailure(
 78                ErrorCode.ALL_PROVIDERS_FAILED,
 79                "All academic search provider pipelines failed",
 80            )
 81
 82        hits: list[PaperSearchHit] = []
 83        for outcome in outcomes:
 84            if isinstance(outcome, list):
 85                hits.extend(outcome)
 86        records = await finalize_paper_hits(
 87            hits,
 88            aggregator=self.aggregator,
 89            resolver=self.resolver,
 90            store=self.store,
 91        )
 92        path = self.result_writer.write_paper_results(
 93            "paper",
 94            records,
 95            request_id=request_id,
 96        )
 97        log_event(
 98            self._logger,
 99            logging.DEBUG,
100            "results_written",
101            kind="paper",
102            path=str(path),
103            results=len(records),
104        )
105        return str(path)
106
107    async def _run_provider(
108        self,
109        provider: AcademicSearchProvider,
110        query: str,
111    ) -> list[PaperSearchHit]:
112        started = self._monotonic()
113        log_event(
114            self._logger,
115            logging.DEBUG,
116            "provider_started",
117            provider=provider.name,
118            stage="paper_search",
119        )
120        try:
121            async with self.quotas.get_academic(provider.name).lease():
122                hits = await provider.search(query)
123            if not isinstance(hits, list):
124                raise TypeError("academic provider search result must be a list")
125            if any(not isinstance(hit, PaperSearchHit) for hit in hits):
126                raise TypeError("academic provider search result contains an invalid item")
127        except asyncio.CancelledError:
128            raise
129        except ExecutionFailure as exc:
130            self._log_provider_failure(provider.name, started, exc)
131            raise
132        except Exception as exc:
133            self._log_provider_failure(provider.name, started, exc)
134            raise ExecutionFailure(
135                ErrorCode.ALL_PROVIDERS_FAILED,
136                f"Academic provider {provider.name} returned invalid data",
137            ) from exc
138        log_event(
139            self._logger,
140            logging.DEBUG,
141            "provider_completed",
142            provider=provider.name,
143            stage="paper_search",
144            results=len(hits),
145            elapsed_ms=elapsed_ms(self._monotonic, started),
146        )
147        return hits
148
149    def _log_provider_failure(
150        self,
151        provider: str,
152        started: float,
153        exc: BaseException,
154    ) -> None:
155        if isinstance(exc, ExecutionFailure) and exc.reason:
156            log_event(
157                self._logger,
158                logging.DEBUG,
159                "provider_failed",
160                provider=provider,
161                stage="paper_search",
162                error_type=type(exc).__name__,
163                elapsed_ms=elapsed_ms(self._monotonic, started),
164                reason=exc.reason,
165            )
166            return
167        log_event(
168            self._logger,
169            logging.DEBUG,
170            "provider_failed",
171            provider=provider,
172            stage="paper_search",
173            error_type=type(exc).__name__,
174            elapsed_ms=elapsed_ms(self._monotonic, started),
175        )
23async def finalize_paper_hits(
24    hits: Sequence[PaperSearchHit],
25    *,
26    aggregator: PaperAggregator,
27    resolver: OAResolver | None,
28    store: URLStore,
29) -> list[PaperRecord]:
30    """Apply the one shared aggregate -> OA enrich -> landing admission policy."""
31
32    records = aggregator.aggregate(hits)
33    enriched = await enrich_paper_records(records, resolver)
34    for record in enriched:
35        admission_abstract = record.abstract.strip() or record.title.strip()
36        store.admit(record.url, admission_abstract)
37    return enriched

Apply the one shared aggregate -> OA enrich -> landing admission policy.

class PaperSearchOrchestrator:
 40class PaperSearchOrchestrator:
 41    def __init__(
 42        self,
 43        *,
 44        providers: Sequence[AcademicSearchProvider],
 45        quotas: ProviderQuotaManager,
 46        aggregator: PaperAggregator,
 47        resolver: OAResolver | None,
 48        store: URLStore,
 49        result_writer: ResultWriter,
 50        logger: logging.Logger | None = None,
 51        monotonic: Callable[[], float] = time.monotonic,
 52    ) -> None:
 53        self.providers = tuple(providers)
 54        self.quotas = quotas
 55        self.aggregator = aggregator
 56        self.resolver = resolver
 57        self.store = store
 58        self.result_writer = result_writer
 59        self._logger = logger or logging.getLogger(__name__)
 60        self._monotonic = monotonic
 61
 62    async def paper_search(self, query: str, *, request_id: str) -> str:
 63        validate_request_id(request_id)
 64        normalized_query = query.strip()
 65        if not normalized_query:
 66            raise InputFailure(ErrorCode.EMPTY_QUERY, "Query must not be empty")
 67        if not self.providers:
 68            raise ExecutionFailure(
 69                ErrorCode.NO_ACADEMIC_SEARCH_PROVIDERS,
 70                "No academic search providers are enabled",
 71            )
 72
 73        outcomes = await asyncio.gather(
 74            *(self._run_provider(provider, normalized_query) for provider in self.providers),
 75            return_exceptions=True,
 76        )
 77        if not any(isinstance(outcome, list) for outcome in outcomes):
 78            raise ExecutionFailure(
 79                ErrorCode.ALL_PROVIDERS_FAILED,
 80                "All academic search provider pipelines failed",
 81            )
 82
 83        hits: list[PaperSearchHit] = []
 84        for outcome in outcomes:
 85            if isinstance(outcome, list):
 86                hits.extend(outcome)
 87        records = await finalize_paper_hits(
 88            hits,
 89            aggregator=self.aggregator,
 90            resolver=self.resolver,
 91            store=self.store,
 92        )
 93        path = self.result_writer.write_paper_results(
 94            "paper",
 95            records,
 96            request_id=request_id,
 97        )
 98        log_event(
 99            self._logger,
100            logging.DEBUG,
101            "results_written",
102            kind="paper",
103            path=str(path),
104            results=len(records),
105        )
106        return str(path)
107
108    async def _run_provider(
109        self,
110        provider: AcademicSearchProvider,
111        query: str,
112    ) -> list[PaperSearchHit]:
113        started = self._monotonic()
114        log_event(
115            self._logger,
116            logging.DEBUG,
117            "provider_started",
118            provider=provider.name,
119            stage="paper_search",
120        )
121        try:
122            async with self.quotas.get_academic(provider.name).lease():
123                hits = await provider.search(query)
124            if not isinstance(hits, list):
125                raise TypeError("academic provider search result must be a list")
126            if any(not isinstance(hit, PaperSearchHit) for hit in hits):
127                raise TypeError("academic provider search result contains an invalid item")
128        except asyncio.CancelledError:
129            raise
130        except ExecutionFailure as exc:
131            self._log_provider_failure(provider.name, started, exc)
132            raise
133        except Exception as exc:
134            self._log_provider_failure(provider.name, started, exc)
135            raise ExecutionFailure(
136                ErrorCode.ALL_PROVIDERS_FAILED,
137                f"Academic provider {provider.name} returned invalid data",
138            ) from exc
139        log_event(
140            self._logger,
141            logging.DEBUG,
142            "provider_completed",
143            provider=provider.name,
144            stage="paper_search",
145            results=len(hits),
146            elapsed_ms=elapsed_ms(self._monotonic, started),
147        )
148        return hits
149
150    def _log_provider_failure(
151        self,
152        provider: str,
153        started: float,
154        exc: BaseException,
155    ) -> None:
156        if isinstance(exc, ExecutionFailure) and exc.reason:
157            log_event(
158                self._logger,
159                logging.DEBUG,
160                "provider_failed",
161                provider=provider,
162                stage="paper_search",
163                error_type=type(exc).__name__,
164                elapsed_ms=elapsed_ms(self._monotonic, started),
165                reason=exc.reason,
166            )
167            return
168        log_event(
169            self._logger,
170            logging.DEBUG,
171            "provider_failed",
172            provider=provider,
173            stage="paper_search",
174            error_type=type(exc).__name__,
175            elapsed_ms=elapsed_ms(self._monotonic, started),
176        )
PaperSearchOrchestrator( *, providers: Sequence[agent_search_gateway.providers.contracts.AcademicSearchProvider], quotas: agent_search_gateway.concurrency.ProviderQuotaManager, aggregator: agent_search_gateway.academic.aggregator.PaperAggregator, resolver: agent_search_gateway.providers.contracts.OAResolver | None, store: agent_search_gateway.url_store.URLStore, result_writer: agent_search_gateway.result_writer.ResultWriter, logger: logging.Logger | None = None, monotonic: Callable[[], float] = <built-in function monotonic>)
41    def __init__(
42        self,
43        *,
44        providers: Sequence[AcademicSearchProvider],
45        quotas: ProviderQuotaManager,
46        aggregator: PaperAggregator,
47        resolver: OAResolver | None,
48        store: URLStore,
49        result_writer: ResultWriter,
50        logger: logging.Logger | None = None,
51        monotonic: Callable[[], float] = time.monotonic,
52    ) -> None:
53        self.providers = tuple(providers)
54        self.quotas = quotas
55        self.aggregator = aggregator
56        self.resolver = resolver
57        self.store = store
58        self.result_writer = result_writer
59        self._logger = logger or logging.getLogger(__name__)
60        self._monotonic = monotonic
providers
quotas
aggregator
resolver
store
result_writer