| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224 |
- #!/usr/bin/env python3
- """Loopback-only OpenAI-compatible server for local acceptance testing.
- This script is deliberately outside the production application. It loads a
- Hugging Face causal language model on CPU and exposes only health and chat
- completion endpoints on 127.0.0.1. Model text is returned unchanged except
- for removing one outer Markdown JSON fence.
- """
- from __future__ import annotations
- import argparse
- import json
- import re
- import time
- import uuid
- from http import HTTPStatus
- from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
- from typing import Any
- MAX_REQUEST_BYTES = 1024 * 1024
- LOOPBACK_HOST = "127.0.0.1"
- def normalize_json_fence(value: str) -> str:
- """Remove one transport-only Markdown fence without changing semantics."""
- text = str(value).strip()
- match = re.fullmatch(
- r"```(?:json)?[ \t]*\r?\n(?P<body>[\s\S]*?)\r?\n```",
- text,
- flags=re.IGNORECASE,
- )
- return match.group("body").strip() if match else text
- class TransformersCPUModel:
- def __init__(self, model_name: str, *, max_tokens: int):
- import torch
- from transformers import AutoModelForCausalLM, AutoTokenizer
- self.model_name = model_name
- self.max_tokens = max_tokens
- self.tokenizer = AutoTokenizer.from_pretrained(
- model_name, local_files_only=True
- )
- self.model = AutoModelForCausalLM.from_pretrained(
- model_name,
- local_files_only=True,
- dtype="auto",
- low_cpu_mem_usage=True,
- ).to("cpu")
- self.model.eval()
- self._torch = torch
- def generate(self, messages: list[dict[str, str]]) -> str:
- prompt = self.tokenizer.apply_chat_template(
- messages,
- tokenize=False,
- add_generation_prompt=True,
- enable_thinking=False,
- )
- inputs = self.tokenizer(
- prompt, return_tensors="pt", add_special_tokens=False
- )
- with self._torch.inference_mode():
- output = self.model.generate(
- **inputs,
- max_new_tokens=self.max_tokens,
- do_sample=False,
- pad_token_id=self.tokenizer.eos_token_id,
- )
- generated = output[0, inputs["input_ids"].shape[1] :]
- return normalize_json_fence(
- self.tokenizer.decode(generated, skip_special_tokens=True)
- )
- class LocalRuleModelServer(ThreadingHTTPServer):
- daemon_threads = True
- def __init__(self, address, model):
- super().__init__(address, LocalRuleModelHandler)
- self.model = model
- class LocalRuleModelHandler(BaseHTTPRequestHandler):
- server: LocalRuleModelServer
- def log_message(self, format: str, *args: Any) -> None:
- print(
- f"{self.address_string()} - {format % args}",
- flush=True,
- )
- def _json(self, status: int, payload: dict[str, Any]) -> None:
- encoded = json.dumps(
- payload, ensure_ascii=False, separators=(",", ":")
- ).encode("utf-8")
- self.send_response(status)
- self.send_header("Content-Type", "application/json; charset=utf-8")
- self.send_header("Content-Length", str(len(encoded)))
- self.send_header("Cache-Control", "no-store")
- self.end_headers()
- self.wfile.write(encoded)
- def do_GET(self) -> None:
- if self.path != "/health":
- self._json(HTTPStatus.NOT_FOUND, {"error": "not found"})
- return
- self._json(
- HTTPStatus.OK,
- {
- "status": "ready",
- "model": self.server.model.model_name,
- "loopback_only": True,
- },
- )
- def do_POST(self) -> None:
- if self.path != "/v1/chat/completions":
- self._json(HTTPStatus.NOT_FOUND, {"error": "not found"})
- return
- try:
- length = int(self.headers.get("Content-Length", "0"))
- except ValueError:
- length = -1
- if length < 1 or length > MAX_REQUEST_BYTES:
- self._json(
- HTTPStatus.REQUEST_ENTITY_TOO_LARGE,
- {"error": {"message": "request size is invalid"}},
- )
- return
- try:
- body = json.loads(self.rfile.read(length))
- messages = body["messages"]
- if (
- not isinstance(messages, list)
- or not messages
- or any(
- not isinstance(item, dict)
- or set(item) != {"role", "content"}
- or item["role"] not in {"system", "user", "assistant"}
- or not isinstance(item["content"], str)
- for item in messages
- )
- ):
- raise ValueError("messages are invalid")
- content = self.server.model.generate(messages)
- except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
- self._json(
- HTTPStatus.BAD_REQUEST,
- {"error": {"message": str(exc) or "request is invalid"}},
- )
- return
- except Exception as exc:
- self._json(
- HTTPStatus.INTERNAL_SERVER_ERROR,
- {
- "error": {
- "message": (
- "local validation model generation failed: "
- f"{type(exc).__name__}"
- )
- }
- },
- )
- return
- self._json(
- HTTPStatus.OK,
- {
- "id": f"chatcmpl-local-{uuid.uuid4().hex}",
- "object": "chat.completion",
- "created": int(time.time()),
- "model": self.server.model.model_name,
- "choices": [
- {
- "index": 0,
- "message": {
- "role": "assistant",
- "content": content,
- },
- "finish_reason": "stop",
- }
- ],
- "usage": {
- "prompt_tokens": 0,
- "completion_tokens": 0,
- "total_tokens": 0,
- },
- },
- )
- def parse_args() -> argparse.Namespace:
- parser = argparse.ArgumentParser(
- description="Loopback-only local OpenAI-compatible validation server"
- )
- parser.add_argument("--host", default=LOOPBACK_HOST)
- parser.add_argument("--port", type=int, required=True)
- parser.add_argument("--model", required=True)
- parser.add_argument("--max-tokens", type=int, default=4096)
- return parser.parse_args()
- def main() -> None:
- args = parse_args()
- if args.host != LOOPBACK_HOST:
- raise SystemExit("local validation server must bind to 127.0.0.1")
- if args.port < 1024 or args.port > 65535:
- raise SystemExit("port must be between 1024 and 65535")
- if args.max_tokens < 512 or args.max_tokens > 8192:
- raise SystemExit("max-tokens must be between 512 and 8192")
- model = TransformersCPUModel(
- args.model, max_tokens=args.max_tokens
- )
- server = LocalRuleModelServer((args.host, args.port), model)
- try:
- server.serve_forever(poll_interval=0.25)
- finally:
- server.server_close()
- if __name__ == "__main__":
- main()
|