"""Card-specific ModelHub XC deployment preflight agent. This service is read-only and only makes claims supported by submitted facts. """ from __future__ import annotations import json import os import re import signal import threading from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any AGENT_NAME = "xc-biren-166m-advisor-agent" AGENT_VERSION = "1.0.0" CARD_VENDOR = "壁仞" CARD_MODEL = "壁砺166M" PORT = int(os.getenv("PORT", "8080")) STRATEGY_ID = os.getenv("STRATEGY_ID", "") MAX_BODY_BYTES = 1_000_000 class ValidationError(ValueError): """Raised when a request field is invalid.""" def _first(payload: dict[str, Any], *names: str) -> Any: for name in names: if name in payload and payload[name] not in (None, ""): return payload[name] return None def _positive_int(value: Any, field: str, default: int) -> int: if value in (None, ""): return default try: number = int(value) except (TypeError, ValueError) as exc: raise ValidationError(f"{field} 必须是整数") from exc if number <= 0: raise ValidationError(f"{field} 必须大于 0") return number def _boolean(value: Any, field: str, default: bool = False) -> bool: if value in (None, ""): return default if isinstance(value, bool): return value if isinstance(value, str): normalized = value.strip().lower() if normalized in {"true", "1", "yes", "on"}: return True if normalized in {"false", "0", "no", "off"}: return False raise ValidationError(f"{field} 必须是布尔值") def _normalize(value: str) -> str: return re.sub(r"[^0-9a-z一-鿿]+", "", value.lower()) def analyze(payload: dict[str, Any]) -> dict[str, Any]: """Return an evidence-labelled preflight for this repository's card.""" if not isinstance(payload, dict): raise ValidationError("请求正文必须是 JSON 对象") model_name = str(_first(payload, "model_name", "model", "model_address") or "") framework = str(_first(payload, "framework") or "") backend = str(_first(payload, "backend", "inference_engine", "engine") or "") sdk_version = str(_first(payload, "sdk_version", "sdk") or "") driver_version = str(_first(payload, "driver_version", "driver") or "") precision = str(_first(payload, "precision", "dtype") or "bf16").lower() cards = _positive_int(_first(payload, "cards", "card_count"), "cards", 1) context_length = _positive_int( _first(payload, "context_length", "max_sequence_length"), "context_length", 4096, ) custom_ops = _boolean(_first(payload, "custom_ops"), "custom_ops") dynamic_shapes = _boolean( _first(payload, "dynamic_shapes"), "dynamic_shapes" ) supplied_hardware = str( _first(payload, "hardware", "device", "target_card", "target_gpu") or "" ) target_tokens = {_normalize(CARD_VENDOR), _normalize(CARD_MODEL)} normalized_hardware = _normalize(supplied_hardware) target_matches = not supplied_hardware or any( token and token in normalized_hardware for token in target_tokens ) environment = { "model_name": model_name, "framework": framework, "backend": backend, "sdk_version": sdk_version, "driver_version": driver_version, } missing_fields = [name for name, value in environment.items() if not value] risks: list[str] = [] if precision in {"int4", "4bit", "int8", "8bit"}: risks.append("量化方案需要实测目标卡后端是否提供对应权重格式与算子内核。") if cards > 1: risks.append("多卡运行需要验证集合通信、进程数、拓扑和并行策略。") if context_length > 32768: risks.append("长上下文会增加 KV Cache 压力,需要按真实并发测峰值内存。") if custom_ops: risks.append("模型包含自定义算子,需要确认编译链、ABI 与目标后端注册情况。") if dynamic_shapes: risks.append("动态形状需要验证图编译缓存、回退路径与重复编译开销。") if not target_matches: verdict = "target_mismatch" risks.insert(0, f"本智能体仅面向 {CARD_VENDOR}|{CARD_MODEL}。") elif missing_fields: verdict = "information_required" else: verdict = "preflight_ready" recommendations = [ f"确认实际目标设备标识为 {CARD_VENDOR}|{CARD_MODEL}。", "记录驱动、SDK、框架和推理后端的完整版本矩阵。", "先用厂商基础样例确认设备可见,再运行模型最小输入。", "依次验证模型加载、首个算子、单请求输出和资源峰值。", "真实支持性结论必须来自目标设备运行日志与健康检查。", ] return { "agent": AGENT_NAME, "version": AGENT_VERSION, "official_target": { "vendor": CARD_VENDOR, "model": CARD_MODEL, "source_scope": "信创模盒模型 X 算力页面", }, "supplied_hardware": supplied_hardware or None, "target_matches": target_matches, "verdict": verdict, "inputs": { **{key: value or None for key, value in environment.items()}, "precision": precision, "cards": cards, "context_length": context_length, "custom_ops": custom_ops, "dynamic_shapes": dynamic_shapes, }, "missing_fields": missing_fields, "risks": risks, "recommendations": recommendations, "disclaimer": "未提供或未实测的硬件规格与兼容性不会被推断为已支持。", } class Handler(BaseHTTPRequestHandler): server_version = "ModelHubCardAdvisor/1.0" def _json(self, payload: dict[str, Any], status: int = 200) -> None: 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 do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract if self.path == "/health": self._json({"status": "ok", "agent": AGENT_NAME, "version": AGENT_VERSION}) return if self.path == "/": self._json( { "name": AGENT_NAME, "version": AGENT_VERSION, "description": f"{CARD_VENDOR} {CARD_MODEL} 专属部署前预检", "official_target": {"vendor": CARD_VENDOR, "model": CARD_MODEL}, "strategy_id_present": bool(STRATEGY_ID), "endpoints": ["GET /health", "POST /analyze", "POST /task"], "external_writes": False, } ) return self._json({"error": "not_found"}, 404) def do_POST(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract if self.path not in {"/analyze", "/task"}: self._json({"error": "not_found"}, 404) return try: length = int(self.headers.get("Content-Length", "0")) if length <= 0: raise ValidationError("请求正文不能为空") if length > MAX_BODY_BYTES: self._json({"error": "payload_too_large"}, 413) return payload = json.loads(self.rfile.read(length)) self._json(analyze(payload)) except (json.JSONDecodeError, UnicodeDecodeError): self._json({"error": "invalid_json", "message": "请求正文必须是有效 JSON"}, 400) except ValidationError as exc: self._json({"error": "validation_error", "message": str(exc)}, 422) def log_message(self, _format: str, *_args: Any) -> None: return STOP = threading.Event() def _handle_signal(signum: int, _frame: Any) -> None: print(f"received signal {signum}; shutting down", flush=True) STOP.set() def main() -> None: signal.signal(signal.SIGTERM, _handle_signal) signal.signal(signal.SIGINT, _handle_signal) server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler) server.timeout = 1 print(f"{AGENT_NAME} {AGENT_VERSION} listening on 0.0.0.0:{PORT}", flush=True) while not STOP.is_set(): server.handle_request() server.server_close() if __name__ == "__main__": main()