diff --git a/backend/package/yuxi/utils/net_security.py b/backend/package/yuxi/utils/net_security.py index 4167fbc3..084f07f7 100644 --- a/backend/package/yuxi/utils/net_security.py +++ b/backend/package/yuxi/utils/net_security.py @@ -76,6 +76,8 @@ class SSRFPolicy(BaseModel): _blocked_networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = PrivateAttr(default_factory=list) _blocked_domains_normalized: set[str] = PrivateAttr(default_factory=set) + _blocked_hostnames_normalized: set[str] = PrivateAttr(default_factory=set) + _blocked_hostname_suffixes_normalized: list[str] = PrivateAttr(default_factory=list) model_config = {"arbitrary_types_allowed": True} @@ -86,6 +88,12 @@ class SSRFPolicy(BaseModel): self._blocked_domains_normalized = { _normalize_hostname(d) for d in self.blocked_domains } + self._blocked_hostnames_normalized = { + _normalize_hostname(h) for h in self.blocked_hostnames + } + self._blocked_hostname_suffixes_normalized = [ + _normalize_hostname(s) for s in self.blocked_hostname_suffixes + ] class SSRFBlockedError(Exception): @@ -109,9 +117,9 @@ def _is_blocked_hostname(hostname: str, policy: SSRFPolicy) -> bool: normalized = _normalize_hostname(hostname) if not normalized: return True # fail-closed - if normalized in policy.blocked_hostnames: + if normalized in policy._blocked_hostnames_normalized: return True - return any(normalized.endswith(suffix) for suffix in policy.blocked_hostname_suffixes) + return any(normalized.endswith(suffix) for suffix in policy._blocked_hostname_suffixes_normalized) def _matches_domain_pattern(hostname: str, pattern: str) -> bool: