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:
Teknium
2026-09-03 13:07:34 -07:00
parent a318098d3c
commit 204c36fd5c
3 changed files with 30 additions and 58 deletions

View File

@@ -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
# ---------------------------------------------------------------------------

View File

@@ -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)

View File

@@ -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)