107 lines
4.6 KiB
Python
107 lines
4.6 KiB
Python
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import logging
|
|
import sys
|
|
import types
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
from common.image_io import ImageLoadError
|
|
|
|
|
|
class _FakeModel:
|
|
def to(self, _device):
|
|
return self
|
|
|
|
def eval(self):
|
|
return self
|
|
|
|
|
|
def _load_clip_module():
|
|
fake_torch = types.ModuleType("torch")
|
|
fake_torch.cuda = types.SimpleNamespace(is_available=lambda: False)
|
|
fake_open_clip = types.ModuleType("open_clip")
|
|
fake_open_clip.create_model_and_transforms = lambda *args, **kwargs: (_FakeModel(), None, lambda image: image)
|
|
fake_open_clip.get_tokenizer = lambda *args, **kwargs: lambda tags: tags
|
|
fake_open_clip.__spec__ = importlib.util.spec_from_loader("open_clip", loader=None)
|
|
fake_torch.__spec__ = importlib.util.spec_from_loader("torch", loader=None)
|
|
|
|
old_modules = {name: sys.modules.get(name) for name in ("torch", "open_clip")}
|
|
sys.modules["torch"] = fake_torch
|
|
sys.modules["open_clip"] = fake_open_clip
|
|
try:
|
|
module_name = "clip_main_observability_test"
|
|
spec = importlib.util.spec_from_file_location(module_name, Path(__file__).parents[1] / "clip" / "main.py")
|
|
module = importlib.util.module_from_spec(spec)
|
|
sys.modules[module_name] = module
|
|
spec.loader.exec_module(module)
|
|
return module
|
|
finally:
|
|
sys.modules.pop("clip_main_observability_test", None)
|
|
for name, previous in old_modules.items():
|
|
if previous is None:
|
|
sys.modules.pop(name, None)
|
|
else:
|
|
sys.modules[name] = previous
|
|
|
|
|
|
class ClipObservabilityTests(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.module = _load_clip_module()
|
|
|
|
def setUp(self):
|
|
self.records = []
|
|
self.handler = logging.Handler()
|
|
self.handler.emit = lambda record: self.records.append(record.getMessage())
|
|
self.module.logger.addHandler(self.handler)
|
|
self.module.logger.setLevel(logging.WARNING)
|
|
|
|
def tearDown(self):
|
|
self.module.logger.removeHandler(self.handler)
|
|
|
|
def assert_failed_embed(self, error):
|
|
url = "https://user:[email protected]/images/art.webp?token=SECRET&sig=VALUE"
|
|
with patch.object(self.module, "fetch_url_bytes", side_effect=error):
|
|
with self.assertRaises(self.module.HTTPException) as raised:
|
|
self.module.embed(self.module.EmbedRequest(url=url))
|
|
self.assertEqual(raised.exception.status_code, 400)
|
|
self.assertEqual(len(self.records), 1)
|
|
line = self.records[0]
|
|
self.assertIn("event='clip_image_fetch_failed'", line)
|
|
for field in ("stage", "error_type", "exception_class", "host", "elapsed_ms"):
|
|
self.assertIn(field, line)
|
|
for secret in ("user:secret", "secret", "?token=", "SECRET", "VALUE", "/images/art.webp"):
|
|
self.assertNotIn(secret, line)
|
|
self.assertNotIn("image-bytes", line)
|
|
self.assertNotIn("Authorization", line)
|
|
|
|
def test_read_timeout_is_logged_and_remains_400(self):
|
|
self.assert_failed_embed(ImageLoadError("read failed", category="read_timeout", stage="fetch", original_exception_class="ReadTimeout"))
|
|
self.assertIn("error_type='read_timeout'", self.records[0])
|
|
|
|
def test_dns_failure_is_logged(self):
|
|
self.assert_failed_embed(ImageLoadError("dns failed", category="dns_failure", stage="validate", original_exception_class="gaierror"))
|
|
self.assertIn("error_type='dns_failure'", self.records[0])
|
|
|
|
def test_invalid_content_type_is_logged(self):
|
|
self.assert_failed_embed(ImageLoadError("not an image", category="invalid_content_type", stage="fetch", original_exception_class="ImageLoadError", metadata={"content_type": "text/html"}))
|
|
self.assertIn("error_type='invalid_content_type'", self.records[0])
|
|
self.assertIn("content_type='text/html'", self.records[0])
|
|
|
|
def test_decode_failure_is_logged(self):
|
|
self.assert_failed_embed(ImageLoadError("decode failed", category="decode_failure", stage="decode", original_exception_class="UnidentifiedImageError"))
|
|
self.assertIn("stage='decode'", self.records[0])
|
|
|
|
def test_success_emits_no_failure_warning(self):
|
|
with patch.object(self.module, "fetch_url_bytes", return_value=b"image-bytes"), patch.object(self.module, "_embed_image_bytes", return_value={"vector": [0.1]}):
|
|
response = self.module.embed(self.module.EmbedRequest(url="https://cdn.example/image.webp"))
|
|
self.assertEqual(response, {"vector": [0.1]})
|
|
self.assertEqual(self.records, [])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|