11
This commit is contained in:
1
app/middleware/__init__.py
Normal file
1
app/middleware/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""ASGI 中间件。"""
|
||||
183
app/middleware/api_crypto.py
Normal file
183
app/middleware/api_crypto.py
Normal file
@@ -0,0 +1,183 @@
|
||||
"""对 /api/v1 请求体解密、响应体加密(需配置 API_AES_KEY)。
|
||||
|
||||
使用纯 ASGI 中间件,避免 BaseHTTPMiddleware 读取 body 后缓存导致路由拿不到明文。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
from app.config import get_settings
|
||||
from app.core.aes_crypto import (
|
||||
ENCRYPTED_HEADER,
|
||||
is_encrypted_envelope,
|
||||
parse_aes_key,
|
||||
unwrap_encrypted,
|
||||
wrap_encrypted,
|
||||
)
|
||||
|
||||
_API_PREFIX = "/api/v1"
|
||||
|
||||
_api_aes_key: bytes | None = None
|
||||
_api_aes_key_initialized = False
|
||||
|
||||
|
||||
def get_api_aes_key() -> bytes | None:
|
||||
print("get_api_aes_key",get_settings().api_aes_key)
|
||||
global _api_aes_key, _api_aes_key_initialized
|
||||
if not _api_aes_key_initialized:
|
||||
_api_aes_key = parse_aes_key(get_settings().api_aes_key)
|
||||
_api_aes_key_initialized = True
|
||||
return _api_aes_key
|
||||
|
||||
|
||||
async def _read_body(receive: Receive) -> bytes:
|
||||
body = b""
|
||||
more = True
|
||||
while more:
|
||||
message = await receive()
|
||||
if message["type"] != "http.request":
|
||||
continue
|
||||
body += message.get("body", b"")
|
||||
more = message.get("more_body", False)
|
||||
return body
|
||||
|
||||
|
||||
def _make_receive(body: bytes) -> Receive:
|
||||
sent = False
|
||||
|
||||
async def receive() -> Message:
|
||||
nonlocal sent
|
||||
if sent:
|
||||
return {"type": "http.disconnect"}
|
||||
sent = True
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
return receive
|
||||
|
||||
|
||||
def _scope_with_content_length(scope: Scope, length: int) -> Scope:
|
||||
new_scope = dict(scope)
|
||||
headers: list[tuple[bytes, bytes]] = []
|
||||
for name, value in scope.get("headers", []):
|
||||
if name.lower() != b"content-length":
|
||||
headers.append((name, value))
|
||||
headers.append((b"content-length", str(length).encode()))
|
||||
new_scope["headers"] = headers
|
||||
return new_scope
|
||||
|
||||
|
||||
def _header_value(scope: Scope, name: str) -> str | None:
|
||||
target = name.lower().encode()
|
||||
for key, value in scope.get("headers", []):
|
||||
if key.lower() == target:
|
||||
return value.decode("latin-1")
|
||||
return None
|
||||
|
||||
|
||||
def _is_json_response(headers: list[tuple[bytes, bytes]]) -> bool:
|
||||
for key, value in headers:
|
||||
if key.lower() == b"content-type":
|
||||
return b"json" in value.lower()
|
||||
return False
|
||||
|
||||
|
||||
def _filter_response_headers(headers: list[tuple[bytes, bytes]]) -> list[tuple[bytes, bytes]]:
|
||||
skip = {b"content-length", b"content-encoding"}
|
||||
return [(k, v) for k, v in headers if k.lower() not in skip]
|
||||
|
||||
|
||||
async def _send_plain_json(send: Send, scope: Scope, status: int, payload: bytes) -> None:
|
||||
await send(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"status": status,
|
||||
"headers": [(b"content-type", b"application/json; charset=utf-8")],
|
||||
}
|
||||
)
|
||||
await send({"type": "http.response.body", "body": payload, "more_body": False})
|
||||
|
||||
|
||||
def _encrypting_send(send: Send, key: bytes) -> Send:
|
||||
status = 200
|
||||
headers: list[tuple[bytes, bytes]] = []
|
||||
chunks: list[bytes] = []
|
||||
started = False
|
||||
|
||||
async def wrapper(message: Message) -> None:
|
||||
nonlocal status, headers, started
|
||||
if message["type"] == "http.response.start":
|
||||
started = True
|
||||
status = message["status"]
|
||||
headers = list(message.get("headers", []))
|
||||
return
|
||||
|
||||
if message["type"] != "http.response.body":
|
||||
await send(message)
|
||||
return
|
||||
|
||||
chunks.append(message.get("body", b""))
|
||||
if message.get("more_body", False):
|
||||
return
|
||||
|
||||
body = b"".join(chunks)
|
||||
out_headers = _filter_response_headers(headers)
|
||||
if _is_json_response(headers):
|
||||
try:
|
||||
body = wrap_encrypted(body, key)
|
||||
except Exception as exc:
|
||||
err = (
|
||||
'{"ok":false,"message":"响应加密失败: '
|
||||
+ str(exc).replace('"', '\\"')
|
||||
+ '"}'
|
||||
).encode("utf-8")
|
||||
await _send_plain_json(send, {}, 500, err)
|
||||
return
|
||||
out_headers.append((ENCRYPTED_HEADER.encode(), b"1"))
|
||||
|
||||
out_headers.append((b"content-length", str(len(body)).encode()))
|
||||
await send({"type": "http.response.start", "status": status, "headers": out_headers})
|
||||
await send({"type": "http.response.body", "body": body, "more_body": False})
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
class ApiCryptoMiddleware:
|
||||
"""纯 ASGI 中间件(勿使用 BaseHTTPMiddleware)。"""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
key = get_api_aes_key()
|
||||
path = scope.get("path", "")
|
||||
|
||||
if key is None or not path.startswith(_API_PREFIX):
|
||||
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
method = scope.get("method", "GET")
|
||||
|
||||
if method in ("POST", "PUT", "PATCH"):
|
||||
body = await _read_body(receive)
|
||||
if (
|
||||
body
|
||||
and _header_value(scope, ENCRYPTED_HEADER) == "1"
|
||||
and is_encrypted_envelope(body)
|
||||
):
|
||||
try:
|
||||
body = unwrap_encrypted(body, key)
|
||||
except Exception as exc:
|
||||
err = f'{{"ok":false,"message":"请求解密失败: {exc}"}}'.encode("utf-8")
|
||||
await _send_plain_json(send, scope, 400, err)
|
||||
return
|
||||
|
||||
receive = _make_receive(body)
|
||||
scope = _scope_with_content_length(scope, len(body))
|
||||
|
||||
await self.app(scope, receive, _encrypting_send(send, key))
|
||||
Reference in New Issue
Block a user