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:secret@cdn.example/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()