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 )
async def
finalize_paper_hits( hits: Sequence[agent_search_gateway.providers.contracts.PaperSearchHit], *, aggregator: agent_search_gateway.academic.aggregator.PaperAggregator, resolver: agent_search_gateway.providers.contracts.OAResolver | None, store: agent_search_gateway.url_store.URLStore) -> list[agent_search_gateway.models.PaperRecord]:
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
async def
paper_search(self, query: str, *, request_id: str) -> str:
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)