mock_api.py 8.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220
  1. """API_PULL 对端 Mock(任务书 P1-D)。
  2. 契约严格对齐 MdpApiPullExecutor.cs 与既有 doc/db/mdp/mock_api/mock_server.py:
  3. - GET 单次请求,不分页;响应数组位于 data.list(可切换 responsePath);
  4. - 每行去重键 bizKey(dedup_key_path);?cursor=<上批最后一行 bizKey> 只返回更大行;
  5. - 鉴权模式 NONE / TOKEN(Bearer) / BASIC / APIKEY,对齐执行器 ApplyAuth 分支;
  6. - 可注入延迟与强制 HTTP 错误(401/403/429/500);
  7. - 端点复用 doc/db/mdp/mock_api/endpoints.json 与 samples/*.json(不迁移、不修改),
  8. 另支持运行时 override:把模拟器适配后的样例(带 SIM 前缀)挂到任意路径。
  9. 每次调用写入内存审计(Header 脱敏),Secret 值以 redact_text 二次遮蔽。
  10. """
  11. from __future__ import annotations
  12. import base64
  13. import json
  14. import sys
  15. import time
  16. from dataclasses import dataclass, field
  17. from pathlib import Path
  18. _SIM_ROOT = Path(__file__).resolve().parents[1]
  19. if str(_SIM_ROOT) not in sys.path:
  20. sys.path.insert(0, str(_SIM_ROOT))
  21. from fastapi import FastAPI, HTTPException, Query, Request # noqa: E402
  22. from catalog.sample_adapters import MOCK_API_DIR # noqa: E402
  23. from clients.redaction import redact_headers # noqa: E402
  24. @dataclass
  25. class MockApiSettings:
  26. auth_mode: str = "TOKEN" # NONE | TOKEN | BASIC | APIKEY
  27. token: str = "sim-mock-token"
  28. basic_user: str = "sim"
  29. basic_password: str = "sim-basic-password"
  30. api_key_header: str = "X-Api-Key"
  31. api_key_value: str = "sim-api-key"
  32. delay_seconds: float = 0.0
  33. force_status: int | None = None # 401/403/429/500 等
  34. response_path: str = "data.list"
  35. audit_limit: int = 200
  36. @dataclass
  37. class MockApiState:
  38. settings: MockApiSettings = field(default_factory=MockApiSettings)
  39. overrides: dict[str, list[dict]] = field(default_factory=dict)
  40. audit: list[dict] = field(default_factory=list)
  41. def secrets(self) -> list[str]:
  42. s = self.settings
  43. return [v for v in (s.token, s.basic_password, s.api_key_value) if v]
  44. def _load_endpoints() -> dict[str, str]:
  45. with (MOCK_API_DIR / "endpoints.json").open("r", encoding="utf-8") as f:
  46. return json.load(f)
  47. def _load_rows(rel_path: str) -> list[dict]:
  48. p = MOCK_API_DIR / rel_path
  49. if not p.exists():
  50. raise HTTPException(status_code=404, detail=f"sample not found: {rel_path}")
  51. with p.open("r", encoding="utf-8") as f:
  52. return json.load(f)
  53. def _build_response(rows: list[dict], response_path: str) -> dict:
  54. parts = [p for p in response_path.split(".") if p]
  55. body: dict = {}
  56. node = body
  57. for part in parts[:-1]:
  58. node[part] = {}
  59. node = node[part]
  60. node[parts[-1]] = rows if parts else rows
  61. body.setdefault("code", 0)
  62. body.setdefault("message", "ok")
  63. return body
  64. def _check_auth(state: MockApiState, request: Request) -> None:
  65. s = state.settings
  66. mode = (s.auth_mode or "NONE").upper()
  67. if mode == "NONE":
  68. return
  69. authorization = request.headers.get("Authorization")
  70. if mode in ("TOKEN", "OAUTH2"):
  71. if authorization != f"Bearer {s.token}":
  72. raise HTTPException(status_code=401, detail="invalid or missing bearer token")
  73. elif mode == "BASIC":
  74. expected = "Basic " + base64.b64encode(f"{s.basic_user}:{s.basic_password}".encode()).decode()
  75. if authorization != expected:
  76. raise HTTPException(status_code=401, detail="invalid or missing basic credentials")
  77. elif mode == "APIKEY":
  78. if request.headers.get(s.api_key_header) != s.api_key_value:
  79. raise HTTPException(status_code=401, detail=f"invalid or missing api key header {s.api_key_header}")
  80. else:
  81. raise HTTPException(status_code=401, detail=f"unsupported auth mode on mock side: {mode}")
  82. def _audit(state: MockApiState, entry: dict) -> None:
  83. state.audit.append(entry)
  84. if len(state.audit) > state.settings.audit_limit:
  85. del state.audit[: len(state.audit) - state.settings.audit_limit]
  86. def create_app(state: MockApiState | None = None) -> FastAPI:
  87. state = state or MockApiState()
  88. app = FastAPI(title="Ai-DOP Integration Simulator - Mock API (API_PULL)", version="1.0.0")
  89. app.state.sim = state
  90. @app.get("/health")
  91. def health():
  92. return {"status": "UP", "authMode": state.settings.auth_mode}
  93. @app.get("/__endpoints")
  94. def list_endpoints():
  95. return {
  96. "authMode": state.settings.auth_mode,
  97. "responsePath": state.settings.response_path,
  98. "endpoints": _load_endpoints(),
  99. "overrides": sorted(state.overrides.keys()),
  100. }
  101. @app.get("/__control/settings")
  102. def get_settings():
  103. s = state.settings
  104. # 不回显任何凭据值
  105. return {
  106. "authMode": s.auth_mode, "delaySeconds": s.delay_seconds,
  107. "forceStatus": s.force_status, "responsePath": s.response_path,
  108. "apiKeyHeader": s.api_key_header,
  109. }
  110. @app.post("/__control/settings")
  111. def update_settings(payload: dict):
  112. s = state.settings
  113. if "auth_mode" in payload:
  114. mode = str(payload["auth_mode"]).upper()
  115. if mode not in ("NONE", "TOKEN", "BASIC", "APIKEY", "OAUTH2"):
  116. raise HTTPException(status_code=400, detail=f"unknown auth_mode: {mode}")
  117. s.auth_mode = mode
  118. if "token" in payload and payload["token"]:
  119. s.token = str(payload["token"])
  120. if "basic_user" in payload and payload["basic_user"]:
  121. s.basic_user = str(payload["basic_user"])
  122. if "basic_password" in payload and payload["basic_password"]:
  123. s.basic_password = str(payload["basic_password"])
  124. if "api_key_header" in payload and payload["api_key_header"]:
  125. s.api_key_header = str(payload["api_key_header"])
  126. if "api_key_value" in payload and payload["api_key_value"]:
  127. s.api_key_value = str(payload["api_key_value"])
  128. if "delay_seconds" in payload:
  129. s.delay_seconds = max(0.0, float(payload["delay_seconds"]))
  130. if "force_status" in payload:
  131. s.force_status = int(payload["force_status"]) if payload["force_status"] else None
  132. if "response_path" in payload and payload["response_path"]:
  133. s.response_path = str(payload["response_path"])
  134. return get_settings()
  135. @app.post("/__control/overrides")
  136. def set_override(payload: dict):
  137. path = str(payload.get("path") or "")
  138. rows = payload.get("rows")
  139. if not path.startswith("/api/") or not isinstance(rows, list):
  140. raise HTTPException(status_code=400, detail="need {path:/api/..., rows:[...]}")
  141. state.overrides[path] = rows
  142. return {"path": path, "rowCount": len(rows)}
  143. @app.delete("/__control/overrides")
  144. def clear_overrides():
  145. state.overrides.clear()
  146. return {"cleared": True}
  147. @app.get("/__control/audit")
  148. def get_audit():
  149. return {"entries": state.audit[-50:]}
  150. @app.get("/api/{obj:path}")
  151. def get_object(
  152. obj: str,
  153. request: Request,
  154. cursor: str | None = Query(default=None, description="上批最后一行 bizKey(增量游标)"),
  155. ):
  156. started = time.time()
  157. path = request.url.path
  158. if state.settings.delay_seconds > 0:
  159. time.sleep(state.settings.delay_seconds)
  160. try:
  161. if state.settings.force_status:
  162. raise HTTPException(status_code=state.settings.force_status,
  163. detail=f"forced error {state.settings.force_status}")
  164. _check_auth(state, request)
  165. if path in state.overrides:
  166. rows = [dict(r) for r in state.overrides[path]]
  167. else:
  168. endpoints = _load_endpoints()
  169. if path not in endpoints:
  170. raise HTTPException(status_code=404,
  171. detail=f"no mock for {path}; add it to endpoints.json or set an override")
  172. rows = _load_rows(endpoints[path])
  173. rows.sort(key=lambda r: str(r.get("bizKey", "")))
  174. total = len(rows)
  175. if cursor:
  176. rows = [r for r in rows if str(r.get("bizKey", "")) > cursor]
  177. return _build_response(rows, state.settings.response_path)
  178. finally:
  179. _audit(state, {
  180. "time": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
  181. "path": path,
  182. "cursor": cursor,
  183. "authMode": state.settings.auth_mode,
  184. "forceStatus": state.settings.force_status,
  185. "elapsedMs": int((time.time() - started) * 1000),
  186. "headers": redact_headers(dict(request.headers), state.secrets()),
  187. })
  188. return app