local_openai_rule_model_server.py 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224
  1. #!/usr/bin/env python3
  2. """Loopback-only OpenAI-compatible server for local acceptance testing.
  3. This script is deliberately outside the production application. It loads a
  4. Hugging Face causal language model on CPU and exposes only health and chat
  5. completion endpoints on 127.0.0.1. Model text is returned unchanged except
  6. for removing one outer Markdown JSON fence.
  7. """
  8. from __future__ import annotations
  9. import argparse
  10. import json
  11. import re
  12. import time
  13. import uuid
  14. from http import HTTPStatus
  15. from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
  16. from typing import Any
  17. MAX_REQUEST_BYTES = 1024 * 1024
  18. LOOPBACK_HOST = "127.0.0.1"
  19. def normalize_json_fence(value: str) -> str:
  20. """Remove one transport-only Markdown fence without changing semantics."""
  21. text = str(value).strip()
  22. match = re.fullmatch(
  23. r"```(?:json)?[ \t]*\r?\n(?P<body>[\s\S]*?)\r?\n```",
  24. text,
  25. flags=re.IGNORECASE,
  26. )
  27. return match.group("body").strip() if match else text
  28. class TransformersCPUModel:
  29. def __init__(self, model_name: str, *, max_tokens: int):
  30. import torch
  31. from transformers import AutoModelForCausalLM, AutoTokenizer
  32. self.model_name = model_name
  33. self.max_tokens = max_tokens
  34. self.tokenizer = AutoTokenizer.from_pretrained(
  35. model_name, local_files_only=True
  36. )
  37. self.model = AutoModelForCausalLM.from_pretrained(
  38. model_name,
  39. local_files_only=True,
  40. dtype="auto",
  41. low_cpu_mem_usage=True,
  42. ).to("cpu")
  43. self.model.eval()
  44. self._torch = torch
  45. def generate(self, messages: list[dict[str, str]]) -> str:
  46. prompt = self.tokenizer.apply_chat_template(
  47. messages,
  48. tokenize=False,
  49. add_generation_prompt=True,
  50. enable_thinking=False,
  51. )
  52. inputs = self.tokenizer(
  53. prompt, return_tensors="pt", add_special_tokens=False
  54. )
  55. with self._torch.inference_mode():
  56. output = self.model.generate(
  57. **inputs,
  58. max_new_tokens=self.max_tokens,
  59. do_sample=False,
  60. pad_token_id=self.tokenizer.eos_token_id,
  61. )
  62. generated = output[0, inputs["input_ids"].shape[1] :]
  63. return normalize_json_fence(
  64. self.tokenizer.decode(generated, skip_special_tokens=True)
  65. )
  66. class LocalRuleModelServer(ThreadingHTTPServer):
  67. daemon_threads = True
  68. def __init__(self, address, model):
  69. super().__init__(address, LocalRuleModelHandler)
  70. self.model = model
  71. class LocalRuleModelHandler(BaseHTTPRequestHandler):
  72. server: LocalRuleModelServer
  73. def log_message(self, format: str, *args: Any) -> None:
  74. print(
  75. f"{self.address_string()} - {format % args}",
  76. flush=True,
  77. )
  78. def _json(self, status: int, payload: dict[str, Any]) -> None:
  79. encoded = json.dumps(
  80. payload, ensure_ascii=False, separators=(",", ":")
  81. ).encode("utf-8")
  82. self.send_response(status)
  83. self.send_header("Content-Type", "application/json; charset=utf-8")
  84. self.send_header("Content-Length", str(len(encoded)))
  85. self.send_header("Cache-Control", "no-store")
  86. self.end_headers()
  87. self.wfile.write(encoded)
  88. def do_GET(self) -> None:
  89. if self.path != "/health":
  90. self._json(HTTPStatus.NOT_FOUND, {"error": "not found"})
  91. return
  92. self._json(
  93. HTTPStatus.OK,
  94. {
  95. "status": "ready",
  96. "model": self.server.model.model_name,
  97. "loopback_only": True,
  98. },
  99. )
  100. def do_POST(self) -> None:
  101. if self.path != "/v1/chat/completions":
  102. self._json(HTTPStatus.NOT_FOUND, {"error": "not found"})
  103. return
  104. try:
  105. length = int(self.headers.get("Content-Length", "0"))
  106. except ValueError:
  107. length = -1
  108. if length < 1 or length > MAX_REQUEST_BYTES:
  109. self._json(
  110. HTTPStatus.REQUEST_ENTITY_TOO_LARGE,
  111. {"error": {"message": "request size is invalid"}},
  112. )
  113. return
  114. try:
  115. body = json.loads(self.rfile.read(length))
  116. messages = body["messages"]
  117. if (
  118. not isinstance(messages, list)
  119. or not messages
  120. or any(
  121. not isinstance(item, dict)
  122. or set(item) != {"role", "content"}
  123. or item["role"] not in {"system", "user", "assistant"}
  124. or not isinstance(item["content"], str)
  125. for item in messages
  126. )
  127. ):
  128. raise ValueError("messages are invalid")
  129. content = self.server.model.generate(messages)
  130. except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
  131. self._json(
  132. HTTPStatus.BAD_REQUEST,
  133. {"error": {"message": str(exc) or "request is invalid"}},
  134. )
  135. return
  136. except Exception as exc:
  137. self._json(
  138. HTTPStatus.INTERNAL_SERVER_ERROR,
  139. {
  140. "error": {
  141. "message": (
  142. "local validation model generation failed: "
  143. f"{type(exc).__name__}"
  144. )
  145. }
  146. },
  147. )
  148. return
  149. self._json(
  150. HTTPStatus.OK,
  151. {
  152. "id": f"chatcmpl-local-{uuid.uuid4().hex}",
  153. "object": "chat.completion",
  154. "created": int(time.time()),
  155. "model": self.server.model.model_name,
  156. "choices": [
  157. {
  158. "index": 0,
  159. "message": {
  160. "role": "assistant",
  161. "content": content,
  162. },
  163. "finish_reason": "stop",
  164. }
  165. ],
  166. "usage": {
  167. "prompt_tokens": 0,
  168. "completion_tokens": 0,
  169. "total_tokens": 0,
  170. },
  171. },
  172. )
  173. def parse_args() -> argparse.Namespace:
  174. parser = argparse.ArgumentParser(
  175. description="Loopback-only local OpenAI-compatible validation server"
  176. )
  177. parser.add_argument("--host", default=LOOPBACK_HOST)
  178. parser.add_argument("--port", type=int, required=True)
  179. parser.add_argument("--model", required=True)
  180. parser.add_argument("--max-tokens", type=int, default=4096)
  181. return parser.parse_args()
  182. def main() -> None:
  183. args = parse_args()
  184. if args.host != LOOPBACK_HOST:
  185. raise SystemExit("local validation server must bind to 127.0.0.1")
  186. if args.port < 1024 or args.port > 65535:
  187. raise SystemExit("port must be between 1024 and 65535")
  188. if args.max_tokens < 512 or args.max_tokens > 8192:
  189. raise SystemExit("max-tokens must be between 512 and 8192")
  190. model = TransformersCPUModel(
  191. args.model, max_tokens=args.max_tokens
  192. )
  193. server = LocalRuleModelServer((args.host, args.port), model)
  194. try:
  195. server.serve_forever(poll_interval=0.25)
  196. finally:
  197. server.server_close()
  198. if __name__ == "__main__":
  199. main()