61 lines
2.0 KiB
Python
61 lines
2.0 KiB
Python
"""`X-Trace-ID` 响应头契约测试。
|
|||
|
|
|
||
|
|
背景:接口文档把它列为所有响应的必备头,但实现里**完全没有**——客户端拿到错误时无法把
|
||
|
|
请求与服务端日志对上。中间件独立成模块正是为了能这样直接测三条分支。
|
||
|
|
"""
|
||
|
|
|
||
|
|
from starlette.requests import Request
|
||
|
|
from starlette.responses import Response
|
||
|
|
|
||
|
|
from app.api.middleware import TRACE_ID_HEADER, attach_trace_id
|
||
|
|
from app.core.contracts import RequestContext
|
||
|
|
|
||
|
|
|
||
|
|
def build_request(
|
||
|
|
headers: dict[str, str] | None = None, *, context: RequestContext | None = None
|
||
|
|
) -> Request:
|
||
|
|
scope = {
|
||
|
|
"type": "http",
|
||
|
|
"method": "GET",
|
||
|
|
"path": "/",
|
||
|
|
"query_string": b"",
|
||
|
|
"headers": [
|
||
|
|
(key.lower().encode("latin-1"), value.encode("latin-1"))
|
||
|
|
for key, value in (headers or {}).items()
|
||
|
|
],
|
||
|
|
}
|
||
|
|
request = Request(scope)
|
||
|
|
if context is not None:
|
||
|
|
request.state.request_context = context
|
||
|
|
return request
|
||
|
|
|
||
|
|
|
||
|
|
async def passthrough(_request: Request) -> Response:
|
||
|
|
return Response()
|
||
|
|
|
||
|
|
|
||
|
|
async def test_request_context_trace_id_wins_over_header() -> None:
|
||
|
|
context = RequestContext(user_id="1", trace_id="server-trace")
|
||
|
|
request = build_request({"X-Trace-ID": "client-trace"}, context=context)
|
||
|
|
|
||
|
|
response = await attach_trace_id(request, passthrough)
|
||
|
|
|
||
|
|
assert response.headers[TRACE_ID_HEADER] == "server-trace"
|
||
|
|
|
||
|
|
|
||
|
|
async def test_falls_back_to_incoming_header_when_context_missing() -> None:
|
||
|
|
"""认证失败等场景还没建立上下文,此时透传客户端带来的 id 而不是凭空造一个。"""
|
||
|
|
request = build_request({"X-Trace-ID": "client-trace"})
|
||
|
|
|
||
|
|
response = await attach_trace_id(request, passthrough)
|
||
|
|
|
||
|
|
assert response.headers[TRACE_ID_HEADER] == "client-trace"
|
||
|
|
|
||
|
|
|
||
|
|
async def test_no_header_when_neither_context_nor_request_header() -> None:
|
||
|
|
request = build_request()
|
||
|
|
|
||
|
|
response = await attach_trace_id(request, passthrough)
|
||
|
|
|
||
|
|
assert TRACE_ID_HEADER not in response.headers
|