"""ModelHub XC error diagnosis agent. The service performs deterministic, read-only diagnosis of submitted runtime logs. It does not execute commands, access external services, or echo secrets. """ from __future__ import annotations import json import os import re import signal import threading from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any, Pattern AGENT_NAME = "xc-error-diagnosis-agent" AGENT_VERSION = "1.0.0" PORT = int(os.getenv("PORT", "8080")) STRATEGY_ID = os.getenv("STRATEGY_ID", "") MAX_BODY_BYTES = 1_000_000 MAX_LOG_CHARS = 200_000 class ValidationError(ValueError): """Raised when a submitted request cannot be diagnosed safely.""" @dataclass(frozen=True) class Rule: category: str severity: str summary: str patterns: tuple[Pattern[str], ...] recommendations: tuple[str, ...] def _patterns(*values: str) -> tuple[Pattern[str], ...]: return tuple(re.compile(value, re.IGNORECASE) for value in values) RULES = ( Rule( category="unsupported_operator", severity="high", summary="目标后端可能缺少模型所需算子或对应内核。", patterns=_patterns( r"unsupported\s+operator", r"not\s+implemented\s+for", r"could\s+not\s+run\s+['\"]?[\w:]+", r"no\s+kernel\s+(?:is\s+)?registered", r"operator\s+[\w:.]+\s+is\s+not\s+supported", ), recommendations=( "核对目标芯片 SDK 与框架后端的算子支持清单。", "确认模型是否启用了 SDPA、FlashAttention 或自定义算子,并尝试平台支持的实现。", "固定框架与 Transformers 版本后,用最小输入复现并记录首个不支持算子。", ), ), Rule( category="out_of_memory", severity="high", summary="设备或主机内存不足,当前模型配置无法完成分配。", patterns=_patterns( r"out\s+of\s+memory", r"memory\s+exhausted", r"alloc(?:ation)?\s+(?:has\s+)?failed", r"cannot\s+allocate\s+memory", r"\bOOM\b", ), recommendations=( "先降低批大小、上下文长度或并发数,记录峰值显存。", "评估更低精度、量化、CPU Offload 或增加卡数。", "确认异常前是否存在残留进程或缓存未释放。", ), ), Rule( category="version_mismatch", severity="high", summary="SDK、驱动、框架或系统运行库之间可能存在版本不匹配。", patterns=_patterns( r"undefined\s+symbol", r"version\s+mismatch", r"requires?\s+.+\s+but\s+found", r"cannot\s+import\s+name", r"GLIBCXX_[0-9.]+\s+not\s+found", ), recommendations=( "记录驱动、芯片 SDK、Python、框架及推理引擎的完整版本矩阵。", "按芯片厂商兼容表选择同一发布周期的组件,避免只升级单个依赖。", "在干净容器中按最小依赖集复现,排除宿主机动态库污染。", ), ), Rule( category="device_runtime", severity="high", summary="目标设备运行时未正确初始化或设备不可见。", patterns=_patterns( r"device\s+(?:is\s+)?unavailable", r"device\s+initiali[sz]ation\s+failed", r"driver\s+version", r"runtime\s+(?:was\s+)?not\s+found", r"no\s+(?:accelerator\s+)?devices?\s+found", ), recommendations=( "确认容器已挂载目标设备并注入厂商运行时。", "检查驱动、固件和 SDK 版本是否匹配。", "先运行厂商自带的设备检测与最小算子样例。", ), ), Rule( category="collective_communication", severity="high", summary="多卡集合通信初始化或同步可能失败。", patterns=_patterns( r"\bNCCL\b.*(?:error|failed|timeout)", r"\bHCCL\b.*(?:error|failed|timeout)", r"collective.*(?:error|failed|timeout)", r"allreduce.*(?:error|failed|timeout)", ), recommendations=( "先验证单卡可运行,再逐步扩展到多卡。", "核对卡间拓扑、通信库版本、网卡选择与进程数配置。", "收集所有 rank 的首个错误,避免只分析后续级联超时。", ), ), Rule( category="missing_dependency", severity="medium", summary="运行环境缺少 Python 模块或动态链接库。", patterns=_patterns( r"ModuleNotFoundError", r"no\s+module\s+named", r"cannot\s+open\s+shared\s+object\s+file", r"shared\s+library\s+not\s+found", ), recommendations=( "根据镜像内实际 Python 解释器核对依赖安装位置。", "确认厂商 SDK 的环境脚本已加载且动态库路径正确。", "将可复现的依赖版本写入镜像构建过程,避免运行时临时安装。", ), ), ) SECRET_PATTERNS = ( re.compile( r"(?i)(authorization\s*[:=]\s*bearer\s+)[A-Za-z0-9._~+/=-]+" ), re.compile( r"(?i)((?:api[_-]?key|access[_-]?token|secret|password|passwd)\s*[:=]\s*)" r"[^\s,;]+" ), ) def redact(text: str) -> tuple[str, bool]: """Remove common credential forms before any submitted text is returned.""" redacted = text for pattern in SECRET_PATTERNS: redacted = pattern.sub(r"\1[REDACTED]", redacted) return redacted, redacted != text 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 _evidence(lines: list[str], rule: Rule) -> list[str]: matches: list[str] = [] for line in lines: compact = line.strip() if not compact: continue if any(pattern.search(compact) for pattern in rule.patterns): matches.append(compact[:500]) if len(matches) == 3: break return matches def analyze(payload: dict[str, Any]) -> dict[str, Any]: """Classify a log using transparent rules and return safe next steps.""" if not isinstance(payload, dict): raise ValidationError("请求正文必须是 JSON 对象") raw_log = _first(payload, "error_log", "log", "logs", "traceback") if raw_log is None: raise ValidationError("error_log 不能为空") if not isinstance(raw_log, str): raise ValidationError("error_log 必须是字符串") if not raw_log.strip(): raise ValidationError("error_log 不能为空") if len(raw_log) > MAX_LOG_CHARS: raise ValidationError(f"error_log 不能超过 {MAX_LOG_CHARS} 个字符") safe_log, redaction_applied = redact(raw_log) lines = safe_log.splitlines() findings: list[dict[str, Any]] = [] recommendations: list[str] = [] for rule in RULES: evidence = _evidence(lines, rule) if not evidence: continue findings.append( { "category": rule.category, "severity": rule.severity, "summary": rule.summary, "evidence": evidence, } ) for recommendation in rule.recommendations: if recommendation not in recommendations: recommendations.append(recommendation) if findings: verdict = "matched_known_failure_patterns" else: verdict = "unknown_pattern" recommendations.extend( [ "提供完整 traceback 中最早出现的异常,而不是只提供最后一行。", "补充模型名、目标芯片、驱动与 SDK、框架和推理引擎版本。", "缩小到最小可复现输入,并确认同一环境下厂商示例是否正常。", ] ) environment = { "hardware": str(_first(payload, "hardware", "device", "target_gpu") or ""), "sdk_version": str(_first(payload, "sdk_version", "sdk") or ""), "framework": str(_first(payload, "framework") or ""), "framework_version": str(_first(payload, "framework_version") or ""), "inference_engine": str(_first(payload, "inference_engine", "engine") or ""), } missing_fields = [name for name, value in environment.items() if not value] return { "agent": AGENT_NAME, "version": AGENT_VERSION, "verdict": verdict, "findings": findings, "environment": {key: value or None for key, value in environment.items()}, "missing_fields": missing_fields, "recommendations": recommendations, "redaction_applied": redaction_applied, "disclaimer": "这是基于日志特征的预诊断,修复前仍需在目标芯片环境中复现验证。", } class Handler(BaseHTTPRequestHandler): server_version = "ModelHubErrorDiagnosisAgent/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": "国产算力运行日志报错预诊断", "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()