add minimal error diagnosis agent
This commit is contained in:
12
Dockerfile
Normal file
12
Dockerfile
Normal file
@@ -0,0 +1,12 @@
|
||||
FROM modelhubxc-4pd.tencentcloudcr.com/xc_agent_platform/python:3.11-slim
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PORT=8080
|
||||
|
||||
WORKDIR /app
|
||||
COPY main.py /app/main.py
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
CMD ["python", "/app/main.py"]
|
||||
49
README.md
49
README.md
@@ -1,3 +1,50 @@
|
||||
# xc-error-diagnosis-agent
|
||||
|
||||
信创模盒国产算力报错预诊断智能体:识别算子、显存、版本、运行时和集合通信问题,并遮蔽敏感信息。
|
||||
信创模盒的国产算力报错预诊断智能体。用户提交运行日志和可选的环境信息后,它会识别常见的算子不支持、显存不足、版本冲突、设备运行时、集合通信与缺失依赖问题,并给出可验证的下一步建议。
|
||||
|
||||
## 安全边界
|
||||
|
||||
- 只使用透明、确定性的规则分析请求内容,不访问外部网络。
|
||||
- 不执行日志中的命令或代码,不自动修改环境。
|
||||
- 返回证据前会遮蔽常见的令牌、密钥和密码形式。
|
||||
- 不回显完整日志,只返回最多三条与规则匹配的证据。
|
||||
- 结论属于预诊断,不替代目标芯片环境中的真实复现。
|
||||
|
||||
## 平台约束
|
||||
|
||||
- `Dockerfile` 位于仓库根目录。
|
||||
- 容器开放 `8080` 端口。
|
||||
- `GET /health` 返回 HTTP 200。
|
||||
- 从环境变量读取平台注入的 `STRATEGY_ID`,仅报告是否存在,不输出其值。
|
||||
- 处理 `SIGTERM`,在平台 30 秒停机窗口内退出。
|
||||
- 仅使用 Python 标准库,适合平台的轻量资源额度。
|
||||
|
||||
## 接口
|
||||
|
||||
- `GET /health`:存活检查。
|
||||
- `GET /`:智能体元信息。
|
||||
- `POST /analyze`:日志预诊断。
|
||||
- `POST /task`:与 `/analyze` 相同的兼容入口。
|
||||
|
||||
示例:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8080/analyze \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"error_log": "RuntimeError: unsupported operator aten::_scaled_dot_product_attention",
|
||||
"hardware": "目标国产算力卡",
|
||||
"sdk_version": "1.0",
|
||||
"framework": "PyTorch",
|
||||
"framework_version": "2.3",
|
||||
"inference_engine": "Transformers"
|
||||
}'
|
||||
```
|
||||
|
||||
## 本地验证
|
||||
|
||||
```bash
|
||||
python3 -m unittest -v
|
||||
python3 main.py
|
||||
curl http://localhost:8080/health
|
||||
```
|
||||
|
||||
329
main.py
Normal file
329
main.py
Normal file
@@ -0,0 +1,329 @@
|
||||
"""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()
|
||||
63
test_main.py
Normal file
63
test_main.py
Normal file
@@ -0,0 +1,63 @@
|
||||
import unittest
|
||||
|
||||
from main import ValidationError, analyze
|
||||
|
||||
|
||||
class AnalyzeTests(unittest.TestCase):
|
||||
def test_unsupported_operator_is_identified(self):
|
||||
result = analyze(
|
||||
{
|
||||
"error_log": (
|
||||
"RuntimeError: unsupported operator "
|
||||
"aten::_scaled_dot_product_attention"
|
||||
),
|
||||
"hardware": "目标国产算力卡",
|
||||
"sdk_version": "1.0",
|
||||
"framework": "PyTorch",
|
||||
"framework_version": "2.3",
|
||||
"inference_engine": "Transformers",
|
||||
}
|
||||
)
|
||||
self.assertEqual(result["verdict"], "matched_known_failure_patterns")
|
||||
self.assertEqual(result["findings"][0]["category"], "unsupported_operator")
|
||||
|
||||
def test_out_of_memory_is_identified(self):
|
||||
result = analyze({"error_log": "RuntimeError: device out of memory"})
|
||||
categories = {finding["category"] for finding in result["findings"]}
|
||||
self.assertIn("out_of_memory", categories)
|
||||
|
||||
def test_version_mismatch_is_identified(self):
|
||||
result = analyze(
|
||||
{"error_log": "ImportError: libbackend.so: undefined symbol: xc_runtime"}
|
||||
)
|
||||
categories = {finding["category"] for finding in result["findings"]}
|
||||
self.assertIn("version_mismatch", categories)
|
||||
|
||||
def test_secrets_are_redacted_from_evidence(self):
|
||||
result = analyze(
|
||||
{
|
||||
"error_log": (
|
||||
"Authorization: Bearer secret-token-123\n"
|
||||
"RuntimeError: unsupported operator aten::example"
|
||||
)
|
||||
}
|
||||
)
|
||||
evidence = "\n".join(
|
||||
line for finding in result["findings"] for line in finding["evidence"]
|
||||
)
|
||||
self.assertNotIn("secret-token-123", evidence)
|
||||
self.assertTrue(result["redaction_applied"])
|
||||
|
||||
def test_unknown_pattern_requests_more_context(self):
|
||||
result = analyze({"error_log": "application exited unexpectedly"})
|
||||
self.assertEqual(result["verdict"], "unknown_pattern")
|
||||
self.assertIn("hardware", result["missing_fields"])
|
||||
self.assertGreaterEqual(len(result["recommendations"]), 3)
|
||||
|
||||
def test_empty_log_is_rejected(self):
|
||||
with self.assertRaises(ValidationError):
|
||||
analyze({"error_log": " "})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user