Files
xc-error-diagnosis-agent/main.py

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()