117 lines
5.0 KiB
Python
117 lines
5.0 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
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)
|
|
|
|
|
|
class QdrantPointSearchTests(unittest.TestCase):
|
|
def test_point_search_retrieves_vector_then_queries(self):
|
|
point = SimpleNamespace(id=1284, vector=[0.1, 0.2], payload=None)
|
|
result = SimpleNamespace(points=[SimpleNamespace(id=59303, score=0.91, payload={"id": 59303})])
|
|
fake = SimpleNamespace()
|
|
fake.retrieve = lambda **kwargs: [point]
|
|
fake.query_points = lambda **kwargs: result
|
|
|
|
with patch.object(module, "client", fake):
|
|
response = module.search_point(module.SearchPointRequest(id=1284, limit=5))
|
|
|
|
self.assertEqual(response["count"], 1)
|
|
self.assertEqual(response["results"][0]["id"], 59303)
|
|
|
|
def test_missing_point_is_404_without_search(self):
|
|
calls = []
|
|
fake = SimpleNamespace(
|
|
retrieve=lambda **kwargs: calls.append(kwargs) or [],
|
|
query_points=lambda **kwargs: self.fail("query_points must not run"),
|
|
)
|
|
|
|
with patch.object(module, "client", fake):
|
|
with self.assertRaises(module.HTTPException) as raised:
|
|
module.search_point(module.SearchPointRequest(id=1284))
|
|
|
|
self.assertEqual(raised.exception.status_code, 404)
|
|
self.assertEqual(calls[0]["ids"], [1284])
|
|
|
|
def test_point_lookup_does_not_generate_an_id(self):
|
|
self.assertEqual(module._lookup_point_id("1284"), 1284)
|
|
self.assertEqual(module._lookup_point_id("0001284"), 1284)
|
|
with self.assertRaises(module.HTTPException) as raised:
|
|
module._lookup_point_id("not-a-point")
|
|
self.assertEqual(raised.exception.status_code, 400)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|