# retoor import asyncio import logging import re from dataclasses import dataclass, field from typing import AsyncIterator, Literal from typosaurus_sandbox.research.client import RsearchClient, RsearchError from typosaurus_sandbox.research.config import ResearchConfig from typosaurus_sandbox.research.envelopes import ChatResponse, DescribeResponse, SearchResponse from typosaurus_sandbox.research.frontier import QueryFrontier, query_variants_from_result logger = logging.getLogger(__name__) ContentKind = Literal["web", "images", "describe", "chat"] URL_PATTERN = re.compile(r"https?://[^\s<>\"']+") @dataclass(frozen=True) class WorkItem: kind: ContentKind value: str deep: bool = False ai: bool = False @dataclass(frozen=True) class Extraction: urls: tuple[str, ...] = () query_seeds: tuple[tuple[str, str], ...] = () content_texts: tuple[str, ...] = () @dataclass class WorkOutcome: item: WorkItem endpoint: str success: bool cache_hit: bool status_code: int | None = None error: str | None = None urls_found: int = 0 queries_seeded: int = 0 contents_seen: int = 0 @dataclass class PipelineReport: outcomes: list[WorkOutcome] = field(default_factory=list) requests_succeeded: int = 0 requests_failed: int = 0 urls_found: int = 0 queries_seeded: int = 0 contents_seen: int = 0 def _urls_from_text(text: str) -> list[str]: cleaned: list[str] = [] for match in URL_PATTERN.findall(text): cleaned.append(match.rstrip(".,;:!?)]}\"'")) return cleaned def extract_response( item: WorkItem, response: SearchResponse | ChatResponse | DescribeResponse, ) -> Extraction: urls: list[str] = [] query_seeds: list[tuple[str, str]] = [] content_texts: list[str] = [] if isinstance(response, SearchResponse): for result in response.results: if result.url: urls.append(result.url) query_seeds.extend(query_variants_from_result(result)) if result.content: content_texts.append(result.content) if response.ai_response: content_texts.append(response.ai_response) query_seeds.append((response.ai_response, "ai_response")) urls.extend(_urls_from_text(response.ai_response)) if response.deep is not None: for source in response.deep.sources: if source.url: urls.append(source.url) query_seeds.extend(query_variants_from_result(source)) if response.deep.markdown: content_texts.append(response.deep.markdown) urls.extend(_urls_from_text(response.deep.markdown)) elif isinstance(response, ChatResponse): if response.response: content_texts.append(response.response) query_seeds.append((response.response, "chat")) urls.extend(_urls_from_text(response.response)) elif isinstance(response, DescribeResponse): if response.description: content_texts.append(response.description) query_seeds.append((response.description, "describe")) urls.extend(_urls_from_text(response.description)) return Extraction( urls=tuple(dict.fromkeys(urls)), query_seeds=tuple(query_seeds), content_texts=tuple(content_texts), ) def apply_extraction(frontier: QueryFrontier, extraction: Extraction) -> tuple[int, int, int]: new_urls = 0 new_queries = 0 new_contents = 0 for url in extraction.urls: if frontier.register_url(url): new_urls += 1 for text, origin in extraction.query_seeds: if frontier.push_query(text, origin): new_queries += 1 for text in extraction.content_texts: if frontier.register_content(text): new_contents += 1 return new_urls, new_queries, new_contents class ResearchPipeline: def __init__(self, client: RsearchClient, frontier: QueryFrontier, config: ResearchConfig | None = None) -> None: self._client = client self._frontier = frontier self._config = config if config is not None else client.config self._pool_size = max(1, self._config.max_concurrency) self._semaphore = asyncio.Semaphore(self._pool_size) @property def pool_size(self) -> int: return self._pool_size @staticmethod def _endpoint(item: WorkItem) -> str: if item.kind in ("web", "images"): return "/search" if item.kind == "describe": return "/describe" return "/chat" def _probe_cache(self, item: WorkItem) -> bool: if item.kind == "web": return ( self._client.search_cached( item.value, content=True, count=self._config.default_count, deep=item.deep, ai=item.ai, ) is not None ) if item.kind == "images": return self._client.search_cached(item.value, type="images", count=self._config.default_count) is not None if item.kind == "describe": return self._client.describe_cached(item.value) is not None return False async def _fetch(self, item: WorkItem) -> SearchResponse | ChatResponse | DescribeResponse: if item.kind == "web": return await self._client.search( item.value, content=True, count=self._config.default_count, deep=item.deep, ai=item.ai, ) if item.kind == "images": return await self._client.search(item.value, type="images", count=self._config.default_count) if item.kind == "describe": return await self._client.describe(item.value) return await self._client.chat(item.value) async def process(self, item: WorkItem) -> WorkOutcome: async with self._semaphore: return await self._handle(item) async def _handle(self, item: WorkItem) -> WorkOutcome: endpoint = self._endpoint(item) cache_hit = self._probe_cache(item) try: response = await self._fetch(item) except RsearchError as exc: outcome = WorkOutcome( item=item, endpoint=endpoint, success=False, cache_hit=cache_hit, status_code=exc.status_code, error=str(exc), ) logger.error( "request failed endpoint=%s kind=%s target=%r status=%s cache_hit=%s error=%s", endpoint, item.kind, item.value, exc.status_code, cache_hit, exc, ) return outcome if isinstance(response, ChatResponse) and response.cached: cache_hit = True if isinstance(response, SearchResponse) and response.deep is not None and response.deep.cache_hit: cache_hit = True extraction = extract_response(item, response) urls_found, queries_seeded, contents_seen = apply_extraction(self._frontier, extraction) logger.debug( "extraction endpoint=%s kind=%s target=%r urls=%s query_seeds=%d content_texts=%d", endpoint, item.kind, item.value, list(extraction.urls), len(extraction.query_seeds), len(extraction.content_texts), ) outcome = WorkOutcome( item=item, endpoint=endpoint, success=True, cache_hit=cache_hit, urls_found=urls_found, queries_seeded=queries_seeded, contents_seen=contents_seen, ) logger.info( "request done endpoint=%s kind=%s target=%r status=ok cache_hit=%s urls=%d queries=%d contents=%d", endpoint, item.kind, item.value, cache_hit, urls_found, queries_seeded, contents_seen, ) return outcome async def run(self, item_source: AsyncIterator[WorkItem]) -> PipelineReport: logger.info("worker pool size=%d max_concurrency=%d", self._pool_size, self._config.max_concurrency) queue: asyncio.Queue[WorkItem | None] = asyncio.Queue(maxsize=self._pool_size * 4) outcomes: list[WorkOutcome] = [] async def produce() -> None: try: async for item in item_source: await queue.put(item) finally: for _ in range(self._pool_size): await queue.put(None) async def consume() -> None: while True: item = await queue.get() if item is None: return try: outcome = await self.process(item) except Exception as exc: logger.error("pool worker error kind=%s target=%r error=%s", item.kind, item.value, exc) continue outcomes.append(outcome) producer_task = asyncio.create_task(produce()) worker_tasks = [asyncio.create_task(consume()) for _ in range(self._pool_size)] try: await producer_task except Exception as exc: logger.error("item source failed error=%s", exc) await asyncio.gather(*worker_tasks) report = self._build_report(outcomes) logger.info( "pipeline finished requests_succeeded=%d requests_failed=%d urls_found=%d queries_seeded=%d contents_seen=%d", report.requests_succeeded, report.requests_failed, report.urls_found, report.queries_seeded, report.contents_seen, ) return report @staticmethod def _build_report(outcomes: list[WorkOutcome]) -> PipelineReport: report = PipelineReport(outcomes=outcomes) for outcome in outcomes: if outcome.success: report.requests_succeeded += 1 else: report.requests_failed += 1 report.urls_found += outcome.urls_found report.queries_seeded += outcome.queries_seeded report.contents_seen += outcome.contents_seen return report