simplify(compat): vision_tools — drop 4 re-exports + 2 shim-only sync validators, repoint 1 caller; SSRF tests now pin the live async validator
This commit is contained in:
@@ -13,7 +13,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from tools.vision_tools import (
|
||||
_validate_image_url,
|
||||
_validate_image_url_async,
|
||||
_handle_vision_analyze,
|
||||
_determine_mime_type,
|
||||
_image_to_base64_data_url,
|
||||
@@ -43,54 +43,43 @@ _RESOLVES = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _validate_image_url — urlparse-based validation
|
||||
# _validate_image_url_async — shape check + SSRF gate on the live download path
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _validate(url) -> bool:
|
||||
return asyncio.run(_validate_image_url_async(url))
|
||||
|
||||
|
||||
class TestValidateImageUrl:
|
||||
"""Tests for URL validation, including urlparse-based netloc check."""
|
||||
"""URL validation on the live (async) download path: urlparse-based netloc check + SSRF."""
|
||||
|
||||
def test_accepts_valid_http_and_https_urls(self):
|
||||
with patch("tools.url_safety.socket.getaddrinfo", return_value=_RESOLVES):
|
||||
assert _validate_image_url("https://example.com/image.jpg") is True
|
||||
assert _validate_image_url("http://cdn.example.org/photo.png") is True
|
||||
assert _validate("https://example.com/image.jpg") is True
|
||||
assert _validate("http://cdn.example.org/photo.png") is True
|
||||
# CDN endpoints that redirect to images should still pass.
|
||||
assert _validate_image_url("https://cdn.example.com/abcdef123") is True
|
||||
assert _validate_image_url("https://img.example.com/pic?w=200&h=200") is True
|
||||
assert _validate_image_url("http://example.com:8080/image.png") is True
|
||||
assert _validate_image_url("https://example.com/") is True
|
||||
assert _validate("https://cdn.example.com/abcdef123") is True
|
||||
assert _validate("https://img.example.com/pic?w=200&h=200") is True
|
||||
assert _validate("http://example.com:8080/image.png") is True
|
||||
assert _validate("https://example.com/") is True
|
||||
|
||||
def test_localhost_url_blocked_by_ssrf(self):
|
||||
"""localhost URLs are blocked by SSRF protection."""
|
||||
assert _validate_image_url("http://localhost:8080/image.png") is False
|
||||
|
||||
"""localhost / loopback URLs are blocked by SSRF protection."""
|
||||
assert _validate("http://localhost:8080/image.png") is False
|
||||
assert _validate("http://127.0.0.1/image.png") is False
|
||||
|
||||
def test_rejects_malformed_and_non_string_inputs(self):
|
||||
# http:// alone has no network location — urlparse catches this.
|
||||
assert _validate_image_url("http://") is False
|
||||
assert _validate_image_url("https://") is False
|
||||
assert _validate_image_url("http:") is False
|
||||
assert _validate_image_url("") is False
|
||||
assert _validate_image_url(" ") is False
|
||||
assert _validate_image_url(None) is False
|
||||
assert _validate_image_url(12345) is False
|
||||
assert _validate_image_url(True) is False
|
||||
assert _validate_image_url(["https://example.com"]) is False
|
||||
|
||||
def test_async_validator_shares_shape_and_ssrf_gate(self):
|
||||
"""The live download path uses the async validator; it must block
|
||||
localhost / malformed input exactly like the sync one."""
|
||||
from tools.vision_tools import _validate_image_url_async
|
||||
|
||||
async def _run():
|
||||
assert await _validate_image_url_async("http://localhost:8080/image.png") is False
|
||||
assert await _validate_image_url_async("http://127.0.0.1/image.png") is False
|
||||
assert await _validate_image_url_async("http://") is False
|
||||
assert await _validate_image_url_async(None) is False
|
||||
with patch("tools.url_safety.socket.getaddrinfo", return_value=_RESOLVES):
|
||||
assert await _validate_image_url_async("https://cdn.example.com/abcdef123") is True
|
||||
|
||||
asyncio.run(_run())
|
||||
assert _validate("http://") is False
|
||||
assert _validate("https://") is False
|
||||
assert _validate("http:") is False
|
||||
assert _validate("") is False
|
||||
assert _validate(" ") is False
|
||||
assert _validate(None) is False
|
||||
assert _validate(12345) is False
|
||||
assert _validate(True) is False
|
||||
assert _validate(["https://example.com"]) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -262,7 +262,7 @@ def _finalize(
|
||||
) -> ResolvedImage:
|
||||
"""Chokepoint: 50MB ingest cap + type check. Images by magic bytes; video (opt-in) by
|
||||
extension + mp4 sniff — enough because every downstream consumer re-validates."""
|
||||
from tools.vision_tools import _detect_image_mime_type_from_bytes
|
||||
from tools.vision_tools_image_prep import _detect_image_mime_type_from_bytes
|
||||
if len(data) > _MAX_INGEST_BYTES:
|
||||
raise SourceTooLarge("media exceeds size limit", src=src, origin=origin)
|
||||
sniffed = _detect_image_mime_type_from_bytes(data)
|
||||
|
||||
@@ -36,17 +36,13 @@ def _load_auxiliary_client() -> None:
|
||||
from hermes_constants import get_hermes_dir
|
||||
from tools.debug_helpers import DebugSession
|
||||
from tools.website_policy import check_website_access
|
||||
from tools.vision_tools_image_prep import ( # noqa: F401 — re-exported for tests/image_source
|
||||
_ANTHROPIC_SUPPORTED_MEDIA_TYPES,
|
||||
from tools.vision_tools_image_prep import (
|
||||
_VISION_MAX_VALIDATED_AGGREGATE_PIXELS,
|
||||
_VISION_MAX_VALIDATED_FRAME_COUNT,
|
||||
_crop_image_region,
|
||||
_detect_image_mime_type_from_bytes,
|
||||
_determine_mime_type,
|
||||
_image_exceeds_dimension,
|
||||
_normalize_to_supported_image,
|
||||
_rasterize_svg_to_png,
|
||||
_supported_media_types,
|
||||
_validate_raster_image_decodable)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -119,23 +115,10 @@ async def _run_encode_on_cpu_executor(fn, *args, **kwargs):
|
||||
return await loop.run_in_executor(_vision_cpu_executor, functools.partial(fn, *args, **kwargs))
|
||||
|
||||
|
||||
def _image_url_shape_ok(url: str) -> bool:
|
||||
"""HTTP(S) shape check only (scheme, netloc). No DNS. Extension-less CDN URLs pass."""
|
||||
return bool(isinstance(url, str) and url.startswith(("http://", "https://")) and urlparse(url).netloc)
|
||||
|
||||
|
||||
def _validate_image_url(url: str) -> bool:
|
||||
"""Validate image URL for sync callers and tests (SSRF via sync DNS check)."""
|
||||
if not _image_url_shape_ok(url):
|
||||
return False
|
||||
# Block private/internal addresses to prevent SSRF
|
||||
from tools.url_safety import is_safe_url
|
||||
return is_safe_url(url)
|
||||
|
||||
|
||||
async def _validate_image_url_async(url: str) -> bool:
|
||||
"""Shape check + SSRF guard with DNS off the event loop."""
|
||||
if not _image_url_shape_ok(url):
|
||||
"""HTTP(S) shape check (scheme, netloc; extension-less CDN URLs pass) + SSRF guard with DNS
|
||||
off the event loop."""
|
||||
if not (isinstance(url, str) and url.startswith(("http://", "https://")) and urlparse(url).netloc):
|
||||
return False
|
||||
from tools.url_safety import async_is_safe_url
|
||||
return await async_is_safe_url(url)
|
||||
|
||||
Reference in New Issue
Block a user