| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135 |
- #!/usr/bin/env python3
- # -*- coding: utf-8 -*-
- """OpenAI 兼容的 mock 模型服务 —— 只为验证 LlmService 的「真实 HTTP 流式」链路。
- 为什么需要它
- ------------
- 本机常常没有可用的大模型端点(Ollama 没启动、外网被 http_proxy 拦掉),
- 但 LlmService 的核心风险恰恰在 HTTP 这一段(SSE 分片解析、usage 提取、
- 错误码映射),全 mock 的单测覆盖不到。这个 mock 端点让 LlmService 真实走一遍:
- DB 模型配置 -> AgentModelFactory -> HTTP SSE -> 分片拼接 -> LlmResult
- 用法
- ----
- # 终端 1
- python mock-openai-server.py 18999
- # 终端 2
- mvn -pl ai-server test -Dtest=LlmServiceHttpTest -Dmock.llm.port=18999
- 模型名约定(测试用来触发特定分支)
- ----------------------------------
- * 普通模型名 -> SSE 逐片吐 CHUNKS,末帧带 usage
- * mock-unauthorized -> 返回 401,验证「模型报错不能被误报成超时」
- * mock-slow -> 先睡 30s,验证超时分支(504)
- """
- import json
- import sys
- import time
- from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
- # 分片拼接后应为:DuckDB 是一个嵌入式的 OLAP 数据库。
- CHUNKS = ["DuckDB", " 是", "一个", "嵌入式", "的 OLAP", " 数据库", "。"]
- PROMPT_TOKENS = 11
- class Handler(BaseHTTPRequestHandler):
- protocol_version = "HTTP/1.1"
- def log_message(self, fmt, *args):
- sys.stderr.write("[mock-openai] " + (fmt % args) + "\n")
- sys.stderr.flush()
- # ---------- 工具 ----------
- def _read_json(self):
- length = int(self.headers.get("Content-Length") or 0)
- raw = self.rfile.read(length) if length > 0 else b"{}"
- try:
- return json.loads(raw.decode("utf-8") or "{}")
- except Exception:
- return {}
- def _write_json(self, status, payload):
- body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
- self.send_response(status)
- self.send_header("Content-Type", "application/json; charset=utf-8")
- self.send_header("Content-Length", str(len(body)))
- self.end_headers()
- self.wfile.write(body)
- def _chunk(self, data: bytes):
- """写一个 HTTP/1.1 chunked 分片"""
- self.wfile.write(("%X\r\n" % len(data)).encode("ascii"))
- self.wfile.write(data)
- self.wfile.write(b"\r\n")
- # ---------- 路由 ----------
- def do_GET(self):
- """健康检查:Java 测试用它确认端口已就绪"""
- self._write_json(200, {"status": "ok", "chunks": len(CHUNKS)})
- def do_POST(self):
- req = self._read_json()
- model = req.get("model") or "mock-model"
- if model == "mock-unauthorized":
- self._write_json(401, {"error": {
- "message": "Incorrect API key provided",
- "type": "invalid_request_error",
- "code": "invalid_api_key",
- }})
- return
- self.send_response(200)
- self.send_header("Content-Type", "text/event-stream; charset=utf-8")
- self.send_header("Cache-Control", "no-cache")
- self.send_header("Transfer-Encoding", "chunked")
- self.end_headers()
- if model == "mock-slow":
- time.sleep(30)
- base = {
- "id": "chatcmpl-mock",
- "object": "chat.completion.chunk",
- "created": int(time.time()),
- "model": model,
- }
- for piece in CHUNKS:
- event = dict(base, choices=[{
- "index": 0,
- "delta": {"content": piece},
- "finish_reason": None,
- }])
- self._chunk(("data: " + json.dumps(event, ensure_ascii=False) + "\n\n").encode("utf-8"))
- final = dict(
- base,
- choices=[{"index": 0, "delta": {}, "finish_reason": "stop"}],
- usage={
- "prompt_tokens": PROMPT_TOKENS,
- "completion_tokens": len(CHUNKS),
- "total_tokens": PROMPT_TOKENS + len(CHUNKS),
- },
- )
- self._chunk(("data: " + json.dumps(final, ensure_ascii=False) + "\n\n").encode("utf-8"))
- self._chunk(b"data: [DONE]\n\n")
- self.wfile.write(b"0\r\n\r\n")
- self.wfile.flush()
- def main():
- port = int(sys.argv[1]) if len(sys.argv) > 1 else 18999
- server = ThreadingHTTPServer(("127.0.0.1", port), Handler)
- sys.stderr.write("[mock-openai] listening on http://127.0.0.1:%d\n" % port)
- sys.stderr.flush()
- server.serve_forever()
- if __name__ == "__main__":
- main()
|