Edit on GitHub

agent_search_gateway.providers.registry

Registration-order-preserving web provider registry.

 1"""Registration-order-preserving web provider registry."""
 2
 3from collections.abc import Callable
 4from dataclasses import dataclass
 5from typing import Literal
 6
 7from .contracts import ProviderCapabilities
 8
 9WebStage = Literal["search", "fetch"]
10WebProviderFactory = Callable[..., object]
11
12
13@dataclass(frozen=True, slots=True)
14class WebProviderRegistration:
15    name: str
16    capabilities: ProviderCapabilities
17    factory: WebProviderFactory
18    allowed_config_keys: frozenset[str]
19    requires_api_key: bool = True
20
21
22class ProviderRegistry:
23    def __init__(self) -> None:
24        self._registrations: dict[str, WebProviderRegistration] = {}
25
26    def register(self, registration: WebProviderRegistration) -> None:
27        if registration.name in self._registrations:
28            raise ValueError(f"provider already registered: {registration.name}")
29        self._registrations[registration.name] = registration
30
31    def get(self, name: str) -> WebProviderRegistration | None:
32        return self._registrations.get(name)
33
34    def require(self, name: str) -> WebProviderRegistration:
35        registration = self.get(name)
36        if registration is None:
37            raise KeyError(name)
38        return registration
39
40    def capabilities(self, name: str) -> ProviderCapabilities:
41        return self.require(name).capabilities
42
43    def list_in_registration_order(self) -> tuple[WebProviderRegistration, ...]:
44        return tuple(self._registrations.values())
45
46    def for_stage(self, stage: WebStage) -> tuple[WebProviderRegistration, ...]:
47        if stage == "search":
48            return tuple(item for item in self._registrations.values() if item.capabilities.search)
49        return tuple(item for item in self._registrations.values() if item.capabilities.fetch)
WebStage = typing.Literal['search', 'fetch']
WebProviderFactory = collections.abc.Callable[..., object]
@dataclass(frozen=True, slots=True)
class WebProviderRegistration:
14@dataclass(frozen=True, slots=True)
15class WebProviderRegistration:
16    name: str
17    capabilities: ProviderCapabilities
18    factory: WebProviderFactory
19    allowed_config_keys: frozenset[str]
20    requires_api_key: bool = True
WebProviderRegistration( name: str, capabilities: agent_search_gateway.providers.contracts.ProviderCapabilities, factory: Callable[..., object], allowed_config_keys: frozenset[str], requires_api_key: bool = True)
name: str
factory: Callable[..., object]
allowed_config_keys: frozenset[str]
requires_api_key: bool
class ProviderRegistry:
23class ProviderRegistry:
24    def __init__(self) -> None:
25        self._registrations: dict[str, WebProviderRegistration] = {}
26
27    def register(self, registration: WebProviderRegistration) -> None:
28        if registration.name in self._registrations:
29            raise ValueError(f"provider already registered: {registration.name}")
30        self._registrations[registration.name] = registration
31
32    def get(self, name: str) -> WebProviderRegistration | None:
33        return self._registrations.get(name)
34
35    def require(self, name: str) -> WebProviderRegistration:
36        registration = self.get(name)
37        if registration is None:
38            raise KeyError(name)
39        return registration
40
41    def capabilities(self, name: str) -> ProviderCapabilities:
42        return self.require(name).capabilities
43
44    def list_in_registration_order(self) -> tuple[WebProviderRegistration, ...]:
45        return tuple(self._registrations.values())
46
47    def for_stage(self, stage: WebStage) -> tuple[WebProviderRegistration, ...]:
48        if stage == "search":
49            return tuple(item for item in self._registrations.values() if item.capabilities.search)
50        return tuple(item for item in self._registrations.values() if item.capabilities.fetch)
def register( self, registration: WebProviderRegistration) -> None:
27    def register(self, registration: WebProviderRegistration) -> None:
28        if registration.name in self._registrations:
29            raise ValueError(f"provider already registered: {registration.name}")
30        self._registrations[registration.name] = registration
def get( self, name: str) -> WebProviderRegistration | None:
32    def get(self, name: str) -> WebProviderRegistration | None:
33        return self._registrations.get(name)
def require( self, name: str) -> WebProviderRegistration:
35    def require(self, name: str) -> WebProviderRegistration:
36        registration = self.get(name)
37        if registration is None:
38            raise KeyError(name)
39        return registration
def capabilities( self, name: str) -> agent_search_gateway.providers.contracts.ProviderCapabilities:
41    def capabilities(self, name: str) -> ProviderCapabilities:
42        return self.require(name).capabilities
def list_in_registration_order( self) -> tuple[WebProviderRegistration, ...]:
44    def list_in_registration_order(self) -> tuple[WebProviderRegistration, ...]:
45        return tuple(self._registrations.values())
def for_stage( self, stage: Literal['search', 'fetch']) -> tuple[WebProviderRegistration, ...]:
47    def for_stage(self, stage: WebStage) -> tuple[WebProviderRegistration, ...]:
48        if stage == "search":
49            return tuple(item for item in self._registrations.values() if item.capabilities.search)
50        return tuple(item for item in self._registrations.values() if item.capabilities.fetch)