test_b5_cors_headers.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124
  1. from __future__ import annotations
  2. import io
  3. import json
  4. from urllib.parse import quote
  5. from dms.api.v1.blueprint import _merge_exposed_headers
  6. ORIGIN = "http://127.0.0.1:9346"
  7. EXPOSED = {"content-disposition", "x-request-id"}
  8. def _headers(token: str | None = None) -> dict[str, str]:
  9. headers = {"Origin": ORIGIN}
  10. if token:
  11. headers["Authorization"] = f"Bearer {token}"
  12. return headers
  13. def _exposed(response) -> list[str]:
  14. return [
  15. value.strip().lower()
  16. for value in response.headers["Access-Control-Expose-Headers"].split(",")
  17. if value.strip()
  18. ]
  19. def test_dms_json_success_exposes_headers_and_request_id_matches(client):
  20. response = client.get("/api/v1/health", headers=_headers())
  21. payload = response.get_json()
  22. assert response.headers["Access-Control-Allow-Origin"] == ORIGIN
  23. assert EXPOSED <= set(_exposed(response))
  24. assert "*" not in _exposed(response)
  25. assert response.headers["X-Request-Id"] == payload["requestId"]
  26. def test_dms_401_and_routing_404_expose_request_id(client):
  27. unauthorized = client.get("/api/v1/documents", headers=_headers())
  28. assert unauthorized.status_code == 401
  29. assert EXPOSED <= set(_exposed(unauthorized))
  30. assert (
  31. unauthorized.headers["X-Request-Id"]
  32. == unauthorized.get_json()["requestId"]
  33. )
  34. missing = client.get("/api/v1/not-implemented", headers=_headers())
  35. assert missing.status_code == 404
  36. assert EXPOSED <= set(_exposed(missing))
  37. assert missing.headers["X-Request-Id"] == missing.get_json()["requestId"]
  38. def test_expose_header_merge_is_case_insensitive_and_deduplicated():
  39. value = _merge_exposed_headers(
  40. "ETag, x-request-id, etag, CONTENT-DISPOSITION"
  41. )
  42. parts = [item.strip().lower() for item in value.split(",")]
  43. assert parts == ["etag", "x-request-id", "content-disposition"]
  44. assert len(parts) == len(set(parts))
  45. assert "*" not in parts
  46. def test_preview_and_chinese_download_expose_headers_once(
  47. b5_app, b5_client, token_for
  48. ):
  49. from dms.extensions import db
  50. from dms.models import AuditLog, Document
  51. token = token_for("admin")
  52. filename = "中文下载文件.pdf"
  53. created = b5_client.post(
  54. "/api/v1/attachments",
  55. data={
  56. "file": (io.BytesIO(b"%PDF-1.4\nCORS"), filename),
  57. "metadata": json.dumps(
  58. {
  59. "documentName": "B5_C1_TEST_中文下载",
  60. "attachmentType": "POLICY",
  61. "summary": "CORS响应头验证",
  62. "tags": ["CORS"],
  63. },
  64. ensure_ascii=False,
  65. ),
  66. },
  67. headers=_headers(token),
  68. content_type="multipart/form-data",
  69. ).get_json()["data"]
  70. preview = b5_client.get(
  71. f"/api/v1/documents/{created['id']}/preview",
  72. headers=_headers(token),
  73. )
  74. assert preview.status_code == 200
  75. assert EXPOSED <= set(_exposed(preview))
  76. assert preview.headers["Content-Disposition"].startswith("inline")
  77. preview.close()
  78. download = b5_client.get(
  79. f"/api/v1/documents/{created['id']}/download",
  80. headers=_headers(token),
  81. )
  82. disposition = download.headers["Content-Disposition"]
  83. assert download.status_code == 200
  84. assert download.headers["Access-Control-Allow-Origin"] == ORIGIN
  85. assert EXPOSED <= set(_exposed(download))
  86. assert "filename=" in disposition
  87. assert "filename*=UTF-8''" in disposition
  88. assert quote(filename, safe="") in disposition
  89. assert download.headers["X-Request-Id"]
  90. download.close()
  91. with b5_app.app_context():
  92. document = db.session.get(Document, int(created["id"]))
  93. audit_count = db.session.scalar(
  94. db.select(db.func.count(AuditLog.id)).where(
  95. AuditLog.target_id == document.id,
  96. AuditLog.action_type == "DOWNLOAD_DOCUMENT",
  97. )
  98. )
  99. assert document.download_count == 1
  100. assert audit_count == 1
  101. def test_ai_response_does_not_gain_dms_expose_headers(client):
  102. response = client.get("/api/health", headers=_headers())
  103. assert response.status_code == 200
  104. assert response.headers["Access-Control-Allow-Origin"] == ORIGIN
  105. assert "Access-Control-Expose-Headers" not in response.headers