Edit on GitHub

agent_search_gateway.providers.web.parallel

Parallel Search and Extract adapter.

  1"""Parallel Search and Extract adapter."""
  2
  3from collections.abc import Mapping
  4
  5from ...errors import ExecutionFailure
  6from ...observability import SecretValue
  7from ...providers.contracts import KeywordSearchHit, URLFetchCandidate
  8from ...url_normalization import NormalizedURL
  9from .common import (
 10    JsonRequester,
 11    endpoint,
 12    failure,
 13    non_empty_string,
 14    normalized_match,
 15    optional_string,
 16    require_list,
 17    require_object,
 18    require_string,
 19)
 20
 21_VALID_MODES = frozenset({"turbo", "fast", "basic", "advanced"})
 22_FETCH_POLICY_KEYS = frozenset({"max_age_seconds", "timeout_seconds", "disable_cache_fallback"})
 23_MIN_MAX_AGE_SECONDS = 600
 24
 25
 26def _validate_mode(value: str | None) -> str | None:
 27    if value is not None and (not isinstance(value, str) or value not in _VALID_MODES):
 28        raise TypeError("mode must be turbo, fast, basic, advanced, or None")
 29    return value
 30
 31
 32def _validate_fetch_policy(
 33    value: Mapping[str, object] | None,
 34    *,
 35    label: str,
 36) -> dict[str, object] | None:
 37    if value is None:
 38        return None
 39    if not isinstance(value, Mapping):
 40        raise TypeError(f"{label} must be a mapping")
 41    if set(value) - _FETCH_POLICY_KEYS:
 42        raise TypeError(f"{label} contains unsupported fields")
 43
 44    if "max_age_seconds" in value:
 45        max_age_seconds = value["max_age_seconds"]
 46        if (
 47            isinstance(max_age_seconds, bool)
 48            or not isinstance(max_age_seconds, int)
 49            or max_age_seconds < _MIN_MAX_AGE_SECONDS
 50        ):
 51            raise TypeError(f"{label}.max_age_seconds must be an integer >= 600")
 52
 53    if "timeout_seconds" in value:
 54        timeout_seconds = value["timeout_seconds"]
 55        if isinstance(timeout_seconds, bool) or not isinstance(timeout_seconds, (int, float)):
 56            raise TypeError(f"{label}.timeout_seconds must be numeric")
 57
 58    if "disable_cache_fallback" in value:
 59        disable_cache_fallback = value["disable_cache_fallback"]
 60        if not isinstance(disable_cache_fallback, bool):
 61            raise TypeError(f"{label}.disable_cache_fallback must be boolean")
 62
 63    return dict(value)
 64
 65
 66class ParallelAdapter:
 67    def __init__(
 68        self,
 69        *,
 70        name: str,
 71        api_url: str,
 72        secret: SecretValue,
 73        http_executor: JsonRequester,
 74        mode: str | None = None,
 75        search_fetch_policy: Mapping[str, object] | None = None,
 76        extract_fetch_policy: Mapping[str, object] | None = None,
 77    ) -> None:
 78        self.name = name
 79        self._api_url = api_url
 80        self._secret = secret
 81        self._http = http_executor
 82        self._mode = _validate_mode(mode)
 83        self._search_fetch_policy = _validate_fetch_policy(
 84            search_fetch_policy,
 85            label="search_fetch_policy",
 86        )
 87        self._extract_fetch_policy = _validate_fetch_policy(
 88            extract_fetch_policy,
 89            label="extract_fetch_policy",
 90        )
 91
 92    @property
 93    def _headers(self) -> dict[str, str]:
 94        return {"x-api-key": self._secret.reveal()}
 95
 96    async def search(self, query: str) -> list[KeywordSearchHit]:
 97        request_body: dict[str, object] = {"search_queries": [query]}
 98        if self._mode is not None:
 99            request_body["mode"] = self._mode
100        if self._search_fetch_policy is not None:
101            request_body["advanced_settings"] = {"fetch_policy": dict(self._search_fetch_policy)}
102
103        payload = await self._http.request_json(
104            "POST",
105            endpoint(self._api_url, "/v1/search"),
106            stage="search",
107            headers=self._headers,
108            json_body=request_body,
109        )
110        root = require_object(payload, self.name, "search", "response")
111        results = require_list(root.get("results"), self.name, "search", "results")
112        hits: list[KeywordSearchHit] = []
113        for item in results:
114            try:
115                result = require_object(item, self.name, "search", "result")
116                url = non_empty_string(result.get("url"), self.name, "search", "result.url")
117                title = optional_string(result.get("title"), self.name, "search", "result.title")
118                excerpts = require_list(
119                    result.get("excerpts"), self.name, "search", "result.excerpts"
120                )
121                snippet = "\n\n".join(
122                    require_string(excerpt, self.name, "search", "result.excerpts[]")
123                    for excerpt in excerpts
124                )
125                hits.append(KeywordSearchHit(url=url, title=title, snippet=snippet))
126            except ExecutionFailure:
127                continue
128        return hits
129
130    async def fetch(self, url: NormalizedURL) -> URLFetchCandidate:
131        advanced_settings: dict[str, object] = {"full_content": True}
132        if self._extract_fetch_policy is not None:
133            advanced_settings["fetch_policy"] = dict(self._extract_fetch_policy)
134
135        payload = await self._http.request_json(
136            "POST",
137            endpoint(self._api_url, "/v1/extract"),
138            stage="fetch",
139            headers=self._headers,
140            json_body={
141                "urls": [str(url)],
142                "advanced_settings": advanced_settings,
143            },
144        )
145        root = require_object(payload, self.name, "fetch", "response")
146        results = require_list(root.get("results"), self.name, "fetch", "results")
147        errors = require_list(root.get("errors"), self.name, "fetch", "errors")
148        for item in results:
149            result = require_object(item, self.name, "fetch", "result")
150            if normalized_match(result.get("url"), url, self.name, "fetch"):
151                full_content = non_empty_string(
152                    result.get("full_content"), self.name, "fetch", "result.full_content"
153                )
154                return URLFetchCandidate(
155                    raw_content=full_content,
156                    content=full_content,
157                )
158        for item in errors:
159            provider_error = require_object(item, self.name, "fetch", "error")
160            if normalized_match(provider_error.get("url"), url, self.name, "fetch"):
161                raise failure(self.name, "fetch", "provider reported extraction failure")
162        raise failure(self.name, "fetch", "matching extraction result was not returned")
class ParallelAdapter:
 67class ParallelAdapter:
 68    def __init__(
 69        self,
 70        *,
 71        name: str,
 72        api_url: str,
 73        secret: SecretValue,
 74        http_executor: JsonRequester,
 75        mode: str | None = None,
 76        search_fetch_policy: Mapping[str, object] | None = None,
 77        extract_fetch_policy: Mapping[str, object] | None = None,
 78    ) -> None:
 79        self.name = name
 80        self._api_url = api_url
 81        self._secret = secret
 82        self._http = http_executor
 83        self._mode = _validate_mode(mode)
 84        self._search_fetch_policy = _validate_fetch_policy(
 85            search_fetch_policy,
 86            label="search_fetch_policy",
 87        )
 88        self._extract_fetch_policy = _validate_fetch_policy(
 89            extract_fetch_policy,
 90            label="extract_fetch_policy",
 91        )
 92
 93    @property
 94    def _headers(self) -> dict[str, str]:
 95        return {"x-api-key": self._secret.reveal()}
 96
 97    async def search(self, query: str) -> list[KeywordSearchHit]:
 98        request_body: dict[str, object] = {"search_queries": [query]}
 99        if self._mode is not None:
100            request_body["mode"] = self._mode
101        if self._search_fetch_policy is not None:
102            request_body["advanced_settings"] = {"fetch_policy": dict(self._search_fetch_policy)}
103
104        payload = await self._http.request_json(
105            "POST",
106            endpoint(self._api_url, "/v1/search"),
107            stage="search",
108            headers=self._headers,
109            json_body=request_body,
110        )
111        root = require_object(payload, self.name, "search", "response")
112        results = require_list(root.get("results"), self.name, "search", "results")
113        hits: list[KeywordSearchHit] = []
114        for item in results:
115            try:
116                result = require_object(item, self.name, "search", "result")
117                url = non_empty_string(result.get("url"), self.name, "search", "result.url")
118                title = optional_string(result.get("title"), self.name, "search", "result.title")
119                excerpts = require_list(
120                    result.get("excerpts"), self.name, "search", "result.excerpts"
121                )
122                snippet = "\n\n".join(
123                    require_string(excerpt, self.name, "search", "result.excerpts[]")
124                    for excerpt in excerpts
125                )
126                hits.append(KeywordSearchHit(url=url, title=title, snippet=snippet))
127            except ExecutionFailure:
128                continue
129        return hits
130
131    async def fetch(self, url: NormalizedURL) -> URLFetchCandidate:
132        advanced_settings: dict[str, object] = {"full_content": True}
133        if self._extract_fetch_policy is not None:
134            advanced_settings["fetch_policy"] = dict(self._extract_fetch_policy)
135
136        payload = await self._http.request_json(
137            "POST",
138            endpoint(self._api_url, "/v1/extract"),
139            stage="fetch",
140            headers=self._headers,
141            json_body={
142                "urls": [str(url)],
143                "advanced_settings": advanced_settings,
144            },
145        )
146        root = require_object(payload, self.name, "fetch", "response")
147        results = require_list(root.get("results"), self.name, "fetch", "results")
148        errors = require_list(root.get("errors"), self.name, "fetch", "errors")
149        for item in results:
150            result = require_object(item, self.name, "fetch", "result")
151            if normalized_match(result.get("url"), url, self.name, "fetch"):
152                full_content = non_empty_string(
153                    result.get("full_content"), self.name, "fetch", "result.full_content"
154                )
155                return URLFetchCandidate(
156                    raw_content=full_content,
157                    content=full_content,
158                )
159        for item in errors:
160            provider_error = require_object(item, self.name, "fetch", "error")
161            if normalized_match(provider_error.get("url"), url, self.name, "fetch"):
162                raise failure(self.name, "fetch", "provider reported extraction failure")
163        raise failure(self.name, "fetch", "matching extraction result was not returned")
ParallelAdapter( *, name: str, api_url: str, secret: agent_search_gateway.observability.SecretValue, http_executor: agent_search_gateway.providers.web.common.JsonRequester, mode: str | None = None, search_fetch_policy: Mapping[str, object] | None = None, extract_fetch_policy: Mapping[str, object] | None = None)
68    def __init__(
69        self,
70        *,
71        name: str,
72        api_url: str,
73        secret: SecretValue,
74        http_executor: JsonRequester,
75        mode: str | None = None,
76        search_fetch_policy: Mapping[str, object] | None = None,
77        extract_fetch_policy: Mapping[str, object] | None = None,
78    ) -> None:
79        self.name = name
80        self._api_url = api_url
81        self._secret = secret
82        self._http = http_executor
83        self._mode = _validate_mode(mode)
84        self._search_fetch_policy = _validate_fetch_policy(
85            search_fetch_policy,
86            label="search_fetch_policy",
87        )
88        self._extract_fetch_policy = _validate_fetch_policy(
89            extract_fetch_policy,
90            label="extract_fetch_policy",
91        )
name
async def search( self, query: str) -> list[agent_search_gateway.providers.contracts.KeywordSearchHit]:
 97    async def search(self, query: str) -> list[KeywordSearchHit]:
 98        request_body: dict[str, object] = {"search_queries": [query]}
 99        if self._mode is not None:
100            request_body["mode"] = self._mode
101        if self._search_fetch_policy is not None:
102            request_body["advanced_settings"] = {"fetch_policy": dict(self._search_fetch_policy)}
103
104        payload = await self._http.request_json(
105            "POST",
106            endpoint(self._api_url, "/v1/search"),
107            stage="search",
108            headers=self._headers,
109            json_body=request_body,
110        )
111        root = require_object(payload, self.name, "search", "response")
112        results = require_list(root.get("results"), self.name, "search", "results")
113        hits: list[KeywordSearchHit] = []
114        for item in results:
115            try:
116                result = require_object(item, self.name, "search", "result")
117                url = non_empty_string(result.get("url"), self.name, "search", "result.url")
118                title = optional_string(result.get("title"), self.name, "search", "result.title")
119                excerpts = require_list(
120                    result.get("excerpts"), self.name, "search", "result.excerpts"
121                )
122                snippet = "\n\n".join(
123                    require_string(excerpt, self.name, "search", "result.excerpts[]")
124                    for excerpt in excerpts
125                )
126                hits.append(KeywordSearchHit(url=url, title=title, snippet=snippet))
127            except ExecutionFailure:
128                continue
129        return hits
131    async def fetch(self, url: NormalizedURL) -> URLFetchCandidate:
132        advanced_settings: dict[str, object] = {"full_content": True}
133        if self._extract_fetch_policy is not None:
134            advanced_settings["fetch_policy"] = dict(self._extract_fetch_policy)
135
136        payload = await self._http.request_json(
137            "POST",
138            endpoint(self._api_url, "/v1/extract"),
139            stage="fetch",
140            headers=self._headers,
141            json_body={
142                "urls": [str(url)],
143                "advanced_settings": advanced_settings,
144            },
145        )
146        root = require_object(payload, self.name, "fetch", "response")
147        results = require_list(root.get("results"), self.name, "fetch", "results")
148        errors = require_list(root.get("errors"), self.name, "fetch", "errors")
149        for item in results:
150            result = require_object(item, self.name, "fetch", "result")
151            if normalized_match(result.get("url"), url, self.name, "fetch"):
152                full_content = non_empty_string(
153                    result.get("full_content"), self.name, "fetch", "result.full_content"
154                )
155                return URLFetchCandidate(
156                    raw_content=full_content,
157                    content=full_content,
158                )
159        for item in errors:
160            provider_error = require_object(item, self.name, "fetch", "error")
161            if normalized_match(provider_error.get("url"), url, self.name, "fetch"):
162                raise failure(self.name, "fetch", "provider reported extraction failure")
163        raise failure(self.name, "fetch", "matching extraction result was not returned")