Preserve CLIP client errors through qdrant

This commit is contained in:
2026-08-29 09:19:38 +02:00
parent d5faa49c47
commit 32ddf70e2f
4 changed files with 122 additions and 9 deletions
+16 -1
View File
@@ -47,6 +47,21 @@ class EmbedRequest(BaseModel):
pretrained: Optional[str] = None pretrained: Optional[str] = None
_TRANSIENT_IMAGE_ERRORS = {
"dns_failure",
"connect_timeout",
"read_timeout",
"connection_error",
"remote_http_5xx",
"fetch_failure",
}
def _image_load_status(error: ImageLoadError) -> int:
"""Keep input failures as 4xx and surface transient/unknown failures as 5xx."""
return 502 if error.category in _TRANSIENT_IMAGE_ERRORS or error.category == "unknown" else 400
def _log_image_load_failure(error: ImageLoadError, url: str, elapsed_ms: float) -> None: def _log_image_load_failure(error: ImageLoadError, url: str, elapsed_ms: float) -> None:
fields = { fields = {
"event": "clip_image_fetch_failed", "event": "clip_image_fetch_failed",
@@ -167,7 +182,7 @@ def embed(req: EmbedRequest):
return _embed_image_bytes(data, backend=req.backend, model_name=req.model, pretrained=req.pretrained) return _embed_image_bytes(data, backend=req.backend, model_name=req.model, pretrained=req.pretrained)
except ImageLoadError as e: except ImageLoadError as e:
_log_image_load_failure(e, req.url, (time.perf_counter() - started_at) * 1000) _log_image_load_failure(e, req.url, (time.perf_counter() - started_at) * 1000)
raise HTTPException(400, str(e)) raise HTTPException(_image_load_status(e), str(e))
@app.post("/embed/file") @app.post("/embed/file")
+9 -2
View File
@@ -211,6 +211,13 @@ def _log_clip_failure(operation: str, *, elapsed_ms: float, status: int | None =
logger.warning("%s", " ".join(f"{key}={value!r}" for key, value in fields.items())) logger.warning("%s", " ".join(f"{key}={value!r}" for key, value in fields.items()))
def _raise_clip_response_error(response: httpx.Response) -> None:
"""Preserve CLIP client statuses while keeping upstream detail private."""
if 400 <= response.status_code < 500:
raise HTTPException(response.status_code, "CLIP image request rejected.")
raise HTTPException(502, "CLIP service unavailable.")
async def _embed_url(url: str) -> List[float]: async def _embed_url(url: str) -> List[float]:
"""Call the CLIP service to get an image embedding.""" """Call the CLIP service to get an image embedding."""
started_at = time.perf_counter() started_at = time.perf_counter()
@@ -222,7 +229,7 @@ async def _embed_url(url: str) -> List[float]:
raise HTTPException(502, f"CLIP request failed: {str(e)}") raise HTTPException(502, f"CLIP request failed: {str(e)}")
if r.status_code >= 400: if r.status_code >= 400:
_log_clip_failure("embed_url", elapsed_ms=(time.perf_counter() - started_at) * 1000, status=r.status_code, error_type="upstream_http_error", safe_detail=_safe_clip_detail(r)) _log_clip_failure("embed_url", elapsed_ms=(time.perf_counter() - started_at) * 1000, status=r.status_code, error_type="upstream_http_error", safe_detail=_safe_clip_detail(r))
raise HTTPException(502, f"CLIP /embed error: {r.status_code} {r.text[:200]}") _raise_clip_response_error(r)
try: try:
return r.json()["vector"] return r.json()["vector"]
except Exception: except Exception:
@@ -242,7 +249,7 @@ async def _embed_bytes(data: bytes) -> List[float]:
raise HTTPException(502, f"CLIP request failed: {str(e)}") raise HTTPException(502, f"CLIP request failed: {str(e)}")
if r.status_code >= 400: if r.status_code >= 400:
_log_clip_failure("embed_file", elapsed_ms=(time.perf_counter() - started_at) * 1000, status=r.status_code, error_type="upstream_http_error", safe_detail=_safe_clip_detail(r)) _log_clip_failure("embed_file", elapsed_ms=(time.perf_counter() - started_at) * 1000, status=r.status_code, error_type="upstream_http_error", safe_detail=_safe_clip_detail(r))
raise HTTPException(502, f"CLIP /embed/file error: {r.status_code} {r.text[:200]}") _raise_clip_response_error(r)
try: try:
return r.json()["vector"] return r.json()["vector"]
except Exception: except Exception:
+18 -6
View File
@@ -62,12 +62,12 @@ class ClipObservabilityTests(unittest.TestCase):
def tearDown(self): def tearDown(self):
self.module.logger.removeHandler(self.handler) self.module.logger.removeHandler(self.handler)
def assert_failed_embed(self, error): def assert_failed_embed(self, error, expected_status):
url = "https://user:[email protected]/images/art.webp?token=SECRET&sig=VALUE" url = "https://user:[email protected]/images/art.webp?token=SECRET&sig=VALUE"
with patch.object(self.module, "fetch_url_bytes", side_effect=error): with patch.object(self.module, "fetch_url_bytes", side_effect=error):
with self.assertRaises(self.module.HTTPException) as raised: with self.assertRaises(self.module.HTTPException) as raised:
self.module.embed(self.module.EmbedRequest(url=url)) self.module.embed(self.module.EmbedRequest(url=url))
self.assertEqual(raised.exception.status_code, 400) self.assertEqual(raised.exception.status_code, expected_status)
self.assertEqual(len(self.records), 1) self.assertEqual(len(self.records), 1)
line = self.records[0] line = self.records[0]
self.assertIn("event='clip_image_fetch_failed'", line) self.assertIn("event='clip_image_fetch_failed'", line)
@@ -79,20 +79,20 @@ class ClipObservabilityTests(unittest.TestCase):
self.assertNotIn("Authorization", line) self.assertNotIn("Authorization", line)
def test_read_timeout_is_logged_and_remains_400(self): 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.assert_failed_embed(ImageLoadError("read failed", category="read_timeout", stage="fetch", original_exception_class="ReadTimeout"), 502)
self.assertIn("error_type='read_timeout'", self.records[0]) self.assertIn("error_type='read_timeout'", self.records[0])
def test_dns_failure_is_logged(self): def test_dns_failure_is_logged(self):
self.assert_failed_embed(ImageLoadError("dns failed", category="dns_failure", stage="validate", original_exception_class="gaierror")) self.assert_failed_embed(ImageLoadError("dns failed", category="dns_failure", stage="validate", original_exception_class="gaierror"), 502)
self.assertIn("error_type='dns_failure'", self.records[0]) self.assertIn("error_type='dns_failure'", self.records[0])
def test_invalid_content_type_is_logged(self): 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.assert_failed_embed(ImageLoadError("not an image", category="invalid_content_type", stage="fetch", original_exception_class="ImageLoadError", metadata={"content_type": "text/html"}), 400)
self.assertIn("error_type='invalid_content_type'", self.records[0]) self.assertIn("error_type='invalid_content_type'", self.records[0])
self.assertIn("content_type='text/html'", self.records[0]) self.assertIn("content_type='text/html'", self.records[0])
def test_decode_failure_is_logged(self): def test_decode_failure_is_logged(self):
self.assert_failed_embed(ImageLoadError("decode failed", category="decode_failure", stage="decode", original_exception_class="UnidentifiedImageError")) self.assert_failed_embed(ImageLoadError("decode failed", category="decode_failure", stage="decode", original_exception_class="UnidentifiedImageError"), 400)
self.assertIn("stage='decode'", self.records[0]) self.assertIn("stage='decode'", self.records[0])
def test_success_emits_no_failure_warning(self): def test_success_emits_no_failure_warning(self):
@@ -101,6 +101,18 @@ class ClipObservabilityTests(unittest.TestCase):
self.assertEqual(response, {"vector": [0.1]}) self.assertEqual(response, {"vector": [0.1]})
self.assertEqual(self.records, []) self.assertEqual(self.records, [])
def test_permanent_categories_remain_400(self):
for category in ("invalid_scheme", "missing_hostname", "blocked_address", "image_too_large", "remote_http_4xx"):
with self.subTest(category=category):
self.records.clear()
self.assert_failed_embed(ImageLoadError("rejected", category=category), 400)
def test_transient_and_unknown_categories_are_502(self):
for category in ("connect_timeout", "connection_error", "remote_http_5xx", "fetch_failure", "unknown"):
with self.subTest(category=category):
self.records.clear()
self.assert_failed_embed(ImageLoadError("upstream failure", category=category), 502)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+79
View File
@@ -0,0 +1,79 @@
from __future__ import annotations
import asyncio
import logging
import unittest
from unittest.mock import patch
import httpx
import qdrant.main as module
class _FakeClient:
def __init__(self, response=None, error=None):
self.response = response
self.error = error
async def __aenter__(self):
return self
async def __aexit__(self, *_args):
return False
async def post(self, *_args, **_kwargs):
if self.error:
raise self.error
return self.response
class QdrantClipSemanticsTests(unittest.IsolatedAsyncioTestCase):
async def call(self, operation, response=None, error=None):
client = _FakeClient(response=response, error=error)
with patch.object(module.httpx, "AsyncClient", lambda **_kwargs: client):
return await operation()
async def test_clip_4xx_is_preserved(self):
for status in (400, 404, 422):
with self.subTest(status=status):
response = httpx.Response(status, json={"detail": "SECRET BODY"}, request=httpx.Request("POST", "http://clip/embed"))
with self.assertRaises(module.HTTPException) as raised:
await self.call(lambda: module._embed_url("https://cdn.example/image.webp"), response=response)
self.assertEqual(raised.exception.status_code, status)
self.assertEqual(raised.exception.detail, "CLIP image request rejected.")
async def test_clip_5xx_and_transport_remain_502(self):
for status in (500, 502, 504):
with self.subTest(status=status):
response = httpx.Response(status, text="SECRET BODY", request=httpx.Request("POST", "http://clip/embed"))
with self.assertRaises(module.HTTPException) as raised:
await self.call(lambda: module._embed_url("https://cdn.example/image.webp"), response=response)
self.assertEqual(raised.exception.status_code, 502)
with self.assertRaises(module.HTTPException) as raised:
await self.call(lambda: module._embed_url("https://cdn.example/image.webp"), error=httpx.ConnectError("secret transport detail"))
self.assertEqual(raised.exception.status_code, 502)
async def test_file_endpoint_preserves_4xx_and_maps_5xx(self):
for status, expected in ((400, 400), (422, 422), (500, 502)):
with self.subTest(status=status):
response = httpx.Response(status, text="SECRET BODY", request=httpx.Request("POST", "http://clip/embed/file"))
with self.assertRaises(module.HTTPException) as raised:
await self.call(lambda: module._embed_bytes(b"image-bytes"), response=response)
self.assertEqual(raised.exception.status_code, expected)
async def test_success_has_no_failure_warning(self):
records = []
handler = logging.Handler()
handler.emit = lambda record: records.append(record.getMessage())
module.logger.addHandler(handler)
try:
response = httpx.Response(200, json={"vector": [0.1]}, request=httpx.Request("POST", "http://clip/embed"))
result = await self.call(lambda: module._embed_url("https://cdn.example/image.webp"), response=response)
self.assertEqual(result, [0.1])
self.assertEqual(records, [])
finally:
module.logger.removeHandler(handler)
if __name__ == "__main__":
unittest.main()