mock-openai-server.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135
  1. #!/usr/bin/env python3
  2. # -*- coding: utf-8 -*-
  3. """OpenAI 兼容的 mock 模型服务 —— 只为验证 LlmService 的「真实 HTTP 流式」链路。
  4. 为什么需要它
  5. ------------
  6. 本机常常没有可用的大模型端点(Ollama 没启动、外网被 http_proxy 拦掉),
  7. 但 LlmService 的核心风险恰恰在 HTTP 这一段(SSE 分片解析、usage 提取、
  8. 错误码映射),全 mock 的单测覆盖不到。这个 mock 端点让 LlmService 真实走一遍:
  9. DB 模型配置 -> AgentModelFactory -> HTTP SSE -> 分片拼接 -> LlmResult
  10. 用法
  11. ----
  12. # 终端 1
  13. python mock-openai-server.py 18999
  14. # 终端 2
  15. mvn -pl ai-server test -Dtest=LlmServiceHttpTest -Dmock.llm.port=18999
  16. 模型名约定(测试用来触发特定分支)
  17. ----------------------------------
  18. * 普通模型名 -> SSE 逐片吐 CHUNKS,末帧带 usage
  19. * mock-unauthorized -> 返回 401,验证「模型报错不能被误报成超时」
  20. * mock-slow -> 先睡 30s,验证超时分支(504)
  21. """
  22. import json
  23. import sys
  24. import time
  25. from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
  26. # 分片拼接后应为:DuckDB 是一个嵌入式的 OLAP 数据库。
  27. CHUNKS = ["DuckDB", " 是", "一个", "嵌入式", "的 OLAP", " 数据库", "。"]
  28. PROMPT_TOKENS = 11
  29. class Handler(BaseHTTPRequestHandler):
  30. protocol_version = "HTTP/1.1"
  31. def log_message(self, fmt, *args):
  32. sys.stderr.write("[mock-openai] " + (fmt % args) + "\n")
  33. sys.stderr.flush()
  34. # ---------- 工具 ----------
  35. def _read_json(self):
  36. length = int(self.headers.get("Content-Length") or 0)
  37. raw = self.rfile.read(length) if length > 0 else b"{}"
  38. try:
  39. return json.loads(raw.decode("utf-8") or "{}")
  40. except Exception:
  41. return {}
  42. def _write_json(self, status, payload):
  43. body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
  44. self.send_response(status)
  45. self.send_header("Content-Type", "application/json; charset=utf-8")
  46. self.send_header("Content-Length", str(len(body)))
  47. self.end_headers()
  48. self.wfile.write(body)
  49. def _chunk(self, data: bytes):
  50. """写一个 HTTP/1.1 chunked 分片"""
  51. self.wfile.write(("%X\r\n" % len(data)).encode("ascii"))
  52. self.wfile.write(data)
  53. self.wfile.write(b"\r\n")
  54. # ---------- 路由 ----------
  55. def do_GET(self):
  56. """健康检查:Java 测试用它确认端口已就绪"""
  57. self._write_json(200, {"status": "ok", "chunks": len(CHUNKS)})
  58. def do_POST(self):
  59. req = self._read_json()
  60. model = req.get("model") or "mock-model"
  61. if model == "mock-unauthorized":
  62. self._write_json(401, {"error": {
  63. "message": "Incorrect API key provided",
  64. "type": "invalid_request_error",
  65. "code": "invalid_api_key",
  66. }})
  67. return
  68. self.send_response(200)
  69. self.send_header("Content-Type", "text/event-stream; charset=utf-8")
  70. self.send_header("Cache-Control", "no-cache")
  71. self.send_header("Transfer-Encoding", "chunked")
  72. self.end_headers()
  73. if model == "mock-slow":
  74. time.sleep(30)
  75. base = {
  76. "id": "chatcmpl-mock",
  77. "object": "chat.completion.chunk",
  78. "created": int(time.time()),
  79. "model": model,
  80. }
  81. for piece in CHUNKS:
  82. event = dict(base, choices=[{
  83. "index": 0,
  84. "delta": {"content": piece},
  85. "finish_reason": None,
  86. }])
  87. self._chunk(("data: " + json.dumps(event, ensure_ascii=False) + "\n\n").encode("utf-8"))
  88. final = dict(
  89. base,
  90. choices=[{"index": 0, "delta": {}, "finish_reason": "stop"}],
  91. usage={
  92. "prompt_tokens": PROMPT_TOKENS,
  93. "completion_tokens": len(CHUNKS),
  94. "total_tokens": PROMPT_TOKENS + len(CHUNKS),
  95. },
  96. )
  97. self._chunk(("data: " + json.dumps(final, ensure_ascii=False) + "\n\n").encode("utf-8"))
  98. self._chunk(b"data: [DONE]\n\n")
  99. self.wfile.write(b"0\r\n\r\n")
  100. self.wfile.flush()
  101. def main():
  102. port = int(sys.argv[1]) if len(sys.argv) > 1 else 18999
  103. server = ThreadingHTTPServer(("127.0.0.1", port), Handler)
  104. sys.stderr.write("[mock-openai] listening on http://127.0.0.1:%d\n" % port)
  105. sys.stderr.flush()
  106. server.serve_forever()
  107. if __name__ == "__main__":
  108. main()