330 lines
12 KiB
Python
330 lines
12 KiB
Python
"""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()
|