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:
Andrey Vasnetsov
2026-05-07 19:37:59 +02:00
committed by George Panchuk
parent beb9ae5b90
commit 07f19f7dee
6 changed files with 261 additions and 1 deletions

View File

@@ -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

View File

@@ -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

View 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)

View File

@@ -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
View 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

View File

@@ -132,6 +132,7 @@ if __name__ == "__main__":
"QdrantRemote": "AsyncQdrantRemote",
"ApiClient": "AsyncApiClient",
"SyncApis": "AsyncApis",
"rest_headers_middleware": "async_rest_headers_middleware",
},
exclude_methods=[
"__del__",