From 6dde24d29351a459a4d232da22dfefbd2570d013 Mon Sep 17 00:00:00 2001 From: huni Date: Tue, 25 Aug 2026 22:25:20 +0800 Subject: [PATCH] add minimal error diagnosis agent --- Dockerfile | 12 ++ README.md | 49 +++++++- main.py | 329 +++++++++++++++++++++++++++++++++++++++++++++++++++ test_main.py | 63 ++++++++++ 4 files changed, 452 insertions(+), 1 deletion(-) create mode 100644 Dockerfile create mode 100644 main.py create mode 100644 test_main.py diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..3cf955d --- /dev/null +++ b/Dockerfile @@ -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"] diff --git a/README.md b/README.md index 481c6bf..0c15607 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,50 @@ # xc-error-diagnosis-agent -信创模盒国产算力报错预诊断智能体:识别算子、显存、版本、运行时和集合通信问题,并遮蔽敏感信息。 \ No newline at end of file +信创模盒的国产算力报错预诊断智能体。用户提交运行日志和可选的环境信息后,它会识别常见的算子不支持、显存不足、版本冲突、设备运行时、集合通信与缺失依赖问题,并给出可验证的下一步建议。 + +## 安全边界 + +- 只使用透明、确定性的规则分析请求内容,不访问外部网络。 +- 不执行日志中的命令或代码,不自动修改环境。 +- 返回证据前会遮蔽常见的令牌、密钥和密码形式。 +- 不回显完整日志,只返回最多三条与规则匹配的证据。 +- 结论属于预诊断,不替代目标芯片环境中的真实复现。 + +## 平台约束 + +- `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 +``` diff --git a/main.py b/main.py new file mode 100644 index 0000000..8946d00 --- /dev/null +++ b/main.py @@ -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() diff --git a/test_main.py b/test_main.py new file mode 100644 index 0000000..34c72f7 --- /dev/null +++ b/test_main.py @@ -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()