mirror of
https://github.com/qdrant/qdrant-client.git
synced 2026-07-23 11:11:01 -05:00
Per-request custom headers for tracing (#1173)
* implement thread-local header overwrite for setting headers for individual requests * mypy * fix: remove redundant import in init, remove redundant import in qdrant remote --------- Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
This commit is contained in:
committed by
George Panchuk
parent
beb9ae5b90
commit
07f19f7dee
@@ -34,6 +34,7 @@ from qdrant_client.conversions.conversion import (
|
||||
grpc_payload_schema_to_field_type,
|
||||
)
|
||||
from qdrant_client.http import AsyncApiClient, AsyncApis, models
|
||||
from qdrant_client.context_headers import async_rest_headers_middleware
|
||||
from qdrant_client.parallel_processor import ParallelWorkerPool
|
||||
from qdrant_client.uploader.grpc_uploader import GrpcBatchUploader
|
||||
from qdrant_client.uploader.rest_uploader import RestBatchUploader
|
||||
@@ -187,6 +188,7 @@ class AsyncQdrantRemote(AsyncQdrantBase):
|
||||
self.openapi_client: AsyncApis[AsyncApiClient] = AsyncApis(
|
||||
host=self.rest_uri, **self._rest_args
|
||||
)
|
||||
self.openapi_client.client.add_middleware(async_rest_headers_middleware)
|
||||
self._grpc_channel_pool: list[grpc.Channel] = []
|
||||
self._grpc_points_client_pool: list[grpc.PointsStub] | None = None
|
||||
self._grpc_collections_client_pool: list[grpc.CollectionsStub] | None = None
|
||||
|
||||
@@ -6,6 +6,7 @@ import grpc
|
||||
|
||||
from qdrant_client.common.client_exceptions import ResourceExhaustedResponse
|
||||
from qdrant_client.common.client_warnings import show_warning_once
|
||||
from qdrant_client.context_headers import get_context_headers
|
||||
|
||||
|
||||
# type: ignore # noqa: F401
|
||||
@@ -172,6 +173,9 @@ def header_adder_interceptor(
|
||||
else:
|
||||
raise ValueError("Synchronous channel requires synchronous auth token provider.")
|
||||
|
||||
for key, value in get_context_headers().items():
|
||||
metadata.append((key, value))
|
||||
|
||||
client_call_details = _ClientCallDetails(
|
||||
client_call_details.method,
|
||||
client_call_details.timeout,
|
||||
@@ -231,6 +235,9 @@ def header_adder_async_interceptor(
|
||||
token = auth_token_provider()
|
||||
metadata.append(("authorization", f"Bearer {token}"))
|
||||
|
||||
for key, value in get_context_headers().items():
|
||||
metadata.append((key, value))
|
||||
|
||||
client_call_details = client_call_details._replace(metadata=metadata)
|
||||
return client_call_details, request_iterator, process_response
|
||||
|
||||
|
||||
47
qdrant_client/context_headers.py
Normal file
47
qdrant_client/context_headers.py
Normal file
@@ -0,0 +1,47 @@
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from contextvars import ContextVar
|
||||
from typing import AsyncIterator, Awaitable, Callable, Iterator
|
||||
|
||||
from httpx import Request, Response
|
||||
|
||||
_context_headers: ContextVar[dict[str, str]] = ContextVar("_context_headers", default={})
|
||||
|
||||
|
||||
def get_context_headers() -> dict[str, str]:
|
||||
return _context_headers.get()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def headers(extra_headers: dict[str, str]) -> Iterator[None]:
|
||||
current = _context_headers.get()
|
||||
merged = {**current, **extra_headers}
|
||||
token = _context_headers.set(merged)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_context_headers.reset(token)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def async_headers(extra_headers: dict[str, str]) -> AsyncIterator[None]:
|
||||
current = _context_headers.get()
|
||||
merged = {**current, **extra_headers}
|
||||
token = _context_headers.set(merged)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_context_headers.reset(token)
|
||||
|
||||
|
||||
def rest_headers_middleware(request: Request, call_next: Callable[[Request], Response]) -> Response:
|
||||
for key, value in get_context_headers().items():
|
||||
request.headers[key] = value
|
||||
return call_next(request)
|
||||
|
||||
|
||||
async def async_rest_headers_middleware(
|
||||
request: Request, call_next: Callable[[Request], Awaitable[Response]]
|
||||
) -> Response:
|
||||
for key, value in get_context_headers().items():
|
||||
request.headers[key] = value
|
||||
return await call_next(request)
|
||||
@@ -1,5 +1,4 @@
|
||||
import importlib.metadata
|
||||
import logging
|
||||
import math
|
||||
import platform
|
||||
from multiprocessing import get_all_start_methods
|
||||
@@ -36,6 +35,7 @@ from qdrant_client.conversions.conversion import (
|
||||
grpc_payload_schema_to_field_type,
|
||||
)
|
||||
from qdrant_client.http import ApiClient, SyncApis, models
|
||||
from qdrant_client.context_headers import rest_headers_middleware
|
||||
from qdrant_client.parallel_processor import ParallelWorkerPool
|
||||
from qdrant_client.uploader.grpc_uploader import GrpcBatchUploader
|
||||
from qdrant_client.uploader.rest_uploader import RestBatchUploader
|
||||
@@ -235,6 +235,8 @@ class QdrantRemote(QdrantBase):
|
||||
**self._rest_args,
|
||||
)
|
||||
|
||||
self.openapi_client.client.add_middleware(rest_headers_middleware)
|
||||
|
||||
self._grpc_channel_pool: list[grpc.Channel] = []
|
||||
self._grpc_points_client_pool: list[grpc.PointsStub] | None = None
|
||||
self._grpc_collections_client_pool: list[grpc.CollectionsStub] | None = None
|
||||
|
||||
201
tests/test_tracing.py
Normal file
201
tests/test_tracing.py
Normal file
@@ -0,0 +1,201 @@
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from httpx import Request
|
||||
|
||||
from qdrant_client.context_headers import (
|
||||
async_headers,
|
||||
async_rest_headers_middleware,
|
||||
get_context_headers,
|
||||
headers,
|
||||
rest_headers_middleware,
|
||||
)
|
||||
|
||||
|
||||
class TestContextHeaders:
|
||||
def test_default_is_empty(self):
|
||||
assert get_context_headers() == {}
|
||||
|
||||
def test_sync_headers_sets_and_resets(self):
|
||||
assert get_context_headers() == {}
|
||||
with headers({"x-tracing-id": "trace-1"}):
|
||||
assert get_context_headers() == {"x-tracing-id": "trace-1"}
|
||||
assert get_context_headers() == {}
|
||||
|
||||
def test_multiple_headers(self):
|
||||
with headers({"x-tracing-id": "trace-1", "x-request-id": "req-1"}):
|
||||
h = get_context_headers()
|
||||
assert h["x-tracing-id"] == "trace-1"
|
||||
assert h["x-request-id"] == "req-1"
|
||||
assert get_context_headers() == {}
|
||||
|
||||
def test_nested_sync_headers(self):
|
||||
with headers({"x-tracing-id": "outer"}):
|
||||
assert get_context_headers() == {"x-tracing-id": "outer"}
|
||||
with headers({"x-tracing-id": "inner"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "inner"
|
||||
assert get_context_headers()["x-tracing-id"] == "outer"
|
||||
assert get_context_headers() == {}
|
||||
|
||||
def test_nested_headers_merge(self):
|
||||
with headers({"x-tracing-id": "trace-1"}):
|
||||
with headers({"x-request-id": "req-1"}):
|
||||
h = get_context_headers()
|
||||
assert h["x-tracing-id"] == "trace-1"
|
||||
assert h["x-request-id"] == "req-1"
|
||||
assert get_context_headers() == {"x-tracing-id": "trace-1"}
|
||||
|
||||
def test_sync_headers_resets_on_exception(self):
|
||||
with pytest.raises(RuntimeError):
|
||||
with headers({"x-tracing-id": "error-trace"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "error-trace"
|
||||
raise RuntimeError("test")
|
||||
assert get_context_headers() == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_headers_sets_and_resets(self):
|
||||
assert get_context_headers() == {}
|
||||
async with async_headers({"x-tracing-id": "async-trace-1"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "async-trace-1"
|
||||
assert get_context_headers() == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nested_async_headers(self):
|
||||
async with async_headers({"x-tracing-id": "outer"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "outer"
|
||||
async with async_headers({"x-tracing-id": "inner"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "inner"
|
||||
assert get_context_headers()["x-tracing-id"] == "outer"
|
||||
assert get_context_headers() == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_headers_resets_on_exception(self):
|
||||
with pytest.raises(RuntimeError):
|
||||
async with async_headers({"x-tracing-id": "error-trace"}):
|
||||
assert get_context_headers()["x-tracing-id"] == "error-trace"
|
||||
raise RuntimeError("test")
|
||||
assert get_context_headers() == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_tasks_get_isolated_context(self):
|
||||
results = {}
|
||||
|
||||
async def task(name: str, trace_id: str):
|
||||
async with async_headers({"x-tracing-id": trace_id}):
|
||||
await asyncio.sleep(0.01)
|
||||
results[name] = get_context_headers()["x-tracing-id"]
|
||||
|
||||
await asyncio.gather(
|
||||
task("a", "trace-a"),
|
||||
task("b", "trace-b"),
|
||||
)
|
||||
assert results["a"] == "trace-a"
|
||||
assert results["b"] == "trace-b"
|
||||
assert get_context_headers() == {}
|
||||
|
||||
|
||||
class TestHttpMiddleware:
|
||||
def test_sync_middleware_adds_headers(self):
|
||||
request = Request("GET", "http://localhost:6333/collections")
|
||||
call_next = MagicMock(return_value="response")
|
||||
|
||||
# Without context - no extra headers
|
||||
result = rest_headers_middleware(request, call_next)
|
||||
assert "x-tracing-id" not in request.headers
|
||||
assert result == "response"
|
||||
|
||||
# With context - headers added
|
||||
request2 = Request("GET", "http://localhost:6333/collections")
|
||||
with headers({"x-tracing-id": "test-123", "x-request-id": "req-456"}):
|
||||
rest_headers_middleware(request2, call_next)
|
||||
assert request2.headers["x-tracing-id"] == "test-123"
|
||||
assert request2.headers["x-request-id"] == "req-456"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_middleware_adds_headers(self):
|
||||
request = Request("GET", "http://localhost:6333/collections")
|
||||
call_next = AsyncMock(return_value="response")
|
||||
|
||||
# Without context - no extra headers
|
||||
await async_rest_headers_middleware(request, call_next)
|
||||
assert "x-tracing-id" not in request.headers
|
||||
|
||||
# With context - headers added
|
||||
request2 = Request("GET", "http://localhost:6333/collections")
|
||||
async with async_headers({"x-tracing-id": "async-123"}):
|
||||
await async_rest_headers_middleware(request2, call_next)
|
||||
assert request2.headers["x-tracing-id"] == "async-123"
|
||||
|
||||
|
||||
class TestGrpcInterceptor:
|
||||
def test_sync_grpc_context_headers(self):
|
||||
from qdrant_client.connection import header_adder_interceptor
|
||||
|
||||
interceptor = header_adder_interceptor(new_metadata=[("api-key", "test")])
|
||||
intercept_fn = interceptor._fn
|
||||
|
||||
class FakeCallDetails:
|
||||
method = "/qdrant.Points/Search"
|
||||
timeout = 5
|
||||
metadata = None
|
||||
credentials = None
|
||||
|
||||
details = FakeCallDetails()
|
||||
|
||||
# Without context headers
|
||||
new_details, _, _ = intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert "x-tracing-id" not in metadata_dict
|
||||
assert metadata_dict["api-key"] == "test"
|
||||
|
||||
# With context headers
|
||||
with headers({"x-tracing-id": "grpc-trace-123", "x-custom": "value"}):
|
||||
new_details, _, _ = intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert metadata_dict["x-tracing-id"] == "grpc-trace-123"
|
||||
assert metadata_dict["x-custom"] == "value"
|
||||
|
||||
# After context - no extra headers
|
||||
new_details, _, _ = intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert "x-tracing-id" not in metadata_dict
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_grpc_context_headers(self):
|
||||
from qdrant_client.connection import header_adder_async_interceptor
|
||||
|
||||
interceptor = header_adder_async_interceptor(new_metadata=[("api-key", "test")])
|
||||
intercept_fn = interceptor._fn
|
||||
|
||||
class FakeCallDetails:
|
||||
method = "/qdrant.Points/Search"
|
||||
timeout = 5
|
||||
metadata = None
|
||||
credentials = None
|
||||
|
||||
def _replace(self, **kwargs):
|
||||
import copy
|
||||
|
||||
new = copy.copy(self)
|
||||
for k, v in kwargs.items():
|
||||
setattr(new, k, v)
|
||||
return new
|
||||
|
||||
details = FakeCallDetails()
|
||||
|
||||
# Without context headers
|
||||
new_details, _, _ = await intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert "x-tracing-id" not in metadata_dict
|
||||
|
||||
# With context headers
|
||||
async with async_headers({"x-tracing-id": "async-grpc-trace"}):
|
||||
new_details, _, _ = await intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert metadata_dict["x-tracing-id"] == "async-grpc-trace"
|
||||
|
||||
# After context - no extra headers
|
||||
new_details, _, _ = await intercept_fn(details, iter(["request"]), False, False)
|
||||
metadata_dict = dict(new_details.metadata)
|
||||
assert "x-tracing-id" not in metadata_dict
|
||||
@@ -132,6 +132,7 @@ if __name__ == "__main__":
|
||||
"QdrantRemote": "AsyncQdrantRemote",
|
||||
"ApiClient": "AsyncApiClient",
|
||||
"SyncApis": "AsyncApis",
|
||||
"rest_headers_middleware": "async_rest_headers_middleware",
|
||||
},
|
||||
exclude_methods=[
|
||||
"__del__",
|
||||
|
||||
Reference in New Issue
Block a user