From 204c36fd5c94d44c7f5eb8431cacf35db2bfd8fd Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 13:07:34 -0700 Subject: [PATCH] =?UTF-8?q?simplify(compat):=20vision=5Ftools=20=E2=80=94?= =?UTF-8?q?=20drop=204=20re-exports=20+=202=20shim-only=20sync=20validator?= =?UTF-8?q?s,=20repoint=201=20caller;=20SSRF=20tests=20now=20pin=20the=20l?= =?UTF-8?q?ive=20async=20validator?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/tools/test_vision_tools.py | 61 +++++++++++++------------------- tools/image_source.py | 2 +- tools/vision_tools.py | 25 +++---------- 3 files changed, 30 insertions(+), 58 deletions(-) diff --git a/tests/tools/test_vision_tools.py b/tests/tools/test_vision_tools.py index 40c96042a3..d27424c9fb 100644 --- a/tests/tools/test_vision_tools.py +++ b/tests/tools/test_vision_tools.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tools/image_source.py b/tools/image_source.py index b0108e3a4e..a123faf4c3 100644 --- a/tools/image_source.py +++ b/tools/image_source.py @@ -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) diff --git a/tools/vision_tools.py b/tools/vision_tools.py index addb59a700..da2d1637f0 100644 --- a/tools/vision_tools.py +++ b/tools/vision_tools.py @@ -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)