Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 50d9c42e02 |
23
README.md
23
README.md
@@ -4,6 +4,9 @@
|
|||||||
|
|
||||||
## 能做什么
|
## 能做什么
|
||||||
|
|
||||||
|
- 后台定时读取信创模盒公开热门模型,自动发现文本生成候选模型。
|
||||||
|
- 选择较小且与官方 EngineX 支持家族相符的候选,按模型与卡型查询历史任务去重。
|
||||||
|
- 当运行环境同时提供私密令牌、策略 ID、合法卡型枚举与验证配置时,自动提交一条验证任务。
|
||||||
- 校验用户指定的目标卡是否与本智能体的官方卡型一致。
|
- 校验用户指定的目标卡是否与本智能体的官方卡型一致。
|
||||||
- 收集模型、框架、后端、驱动和 SDK 版本,明确缺失信息。
|
- 收集模型、框架、后端、驱动和 SDK 版本,明确缺失信息。
|
||||||
- 根据多卡、量化、长上下文、自定义算子和动态形状提示验证风险。
|
- 根据多卡、量化、长上下文、自定义算子和动态形状提示验证风险。
|
||||||
@@ -13,9 +16,26 @@
|
|||||||
|
|
||||||
- 卡型名称来自信创模盒“模型 X 算力”页面。
|
- 卡型名称来自信创模盒“模型 X 算力”页面。
|
||||||
- 不声明页面未提供的显存、算力、SDK 版本或算子支持结论。
|
- 不声明页面未提供的显存、算力、SDK 版本或算子支持结论。
|
||||||
- 不访问外部网络,不执行命令,不创建或修改适配任务。
|
- 只访问信创模盒公开候选接口和官方任务接口;不执行系统命令。
|
||||||
|
- 无私密 `MODELHUB_XC_TOKEN` 时只扫描、记录候选,不创建任务。
|
||||||
|
- 令牌只从运行环境读取,不写入仓库、HTTP 响应或 stdout。
|
||||||
- 最终兼容性必须以目标设备上的真实日志和测试结果为准。
|
- 最终兼容性必须以目标设备上的真实日志和测试结果为准。
|
||||||
|
|
||||||
|
## 自动工作模式
|
||||||
|
|
||||||
|
- `AUTO_SCAN_ENABLED`:默认 `true`,启动后立即扫描,之后每小时扫描一次。
|
||||||
|
- `SCAN_INTERVAL_SECONDS`:扫描间隔,最小 60 秒,默认 3600 秒。
|
||||||
|
- `AUTO_SUBMIT_ENABLED`:默认 `true`;缺少任一必要配置时自动降级为只读扫描。
|
||||||
|
- `MODELHUB_XC_TOKEN`:调用任务查询/创建接口的私密 `Xc-Token`,不得写进公开代码。
|
||||||
|
- `MODELHUB_TARGET_GPU`:平台内部卡型枚举。只有已由官方任务记录核验的枚举才可配置。
|
||||||
|
- `MODELHUB_CONFIG_PARAMS`:验证任务 YAML;未内置官方配置的卡型必须通过私密部署配置提供。
|
||||||
|
- 每次扫描最多提交一个候选;提交前用 `modelId + gpuType` 精确查询历史,已有任务则跳过。
|
||||||
|
|
||||||
|
## 日志
|
||||||
|
|
||||||
|
- stdout 使用单行 JSON,记录服务启动、扫描、候选、去重、提交结果、健康检查、分析结论和优雅停机。
|
||||||
|
- 仅记录经过清洗的模型标识、结论和必要状态,不输出整个请求体、凭据、URL 查询参数或环境变量值。
|
||||||
|
|
||||||
## 平台约束
|
## 平台约束
|
||||||
|
|
||||||
- 仓库根目录包含 `Dockerfile`,容器开放 `8080`。
|
- 仓库根目录包含 `Dockerfile`,容器开放 `8080`。
|
||||||
@@ -28,6 +48,7 @@
|
|||||||
|
|
||||||
- `GET /health`:存活检查。
|
- `GET /health`:存活检查。
|
||||||
- `GET /`:卡型和智能体元信息。
|
- `GET /`:卡型和智能体元信息。
|
||||||
|
- `GET /scanner`:后台扫描器最近状态,不包含令牌或策略 ID。
|
||||||
- `POST /analyze`:部署前预检。
|
- `POST /analyze`:部署前预检。
|
||||||
- `POST /task`:与 `/analyze` 相同的兼容入口。
|
- `POST /task`:与 `/analyze` 相同的兼容入口。
|
||||||
|
|
||||||
|
|||||||
348
main.py
348
main.py
@@ -1,6 +1,7 @@
|
|||||||
"""Card-specific ModelHub XC deployment preflight agent.
|
"""Card-specific ModelHub XC deployment preflight agent.
|
||||||
|
|
||||||
This service is read-only and only makes claims supported by submitted facts.
|
It scans public candidates and only writes a deduplicated task when private runtime
|
||||||
|
credentials and verified platform configuration are present.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -10,23 +11,279 @@ import os
|
|||||||
import re
|
import re
|
||||||
import signal
|
import signal
|
||||||
import threading
|
import threading
|
||||||
|
from datetime import datetime, timezone
|
||||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
from urllib.parse import quote, urlencode, urlsplit
|
||||||
|
from urllib.request import ProxyHandler, Request, build_opener
|
||||||
|
|
||||||
|
|
||||||
AGENT_NAME = "xc-tianga100-advisor-agent"
|
AGENT_NAME = "xc-tianga100-advisor-agent"
|
||||||
AGENT_VERSION = "1.0.0"
|
AGENT_VERSION = "1.2.0"
|
||||||
CARD_VENDOR = "天数智芯"
|
CARD_VENDOR = "天数智芯"
|
||||||
CARD_MODEL = "天垓100"
|
CARD_MODEL = "天垓100"
|
||||||
|
DEFAULT_TARGET_GPU = "Iluvatar_bi-100"
|
||||||
PORT = int(os.getenv("PORT", "8080"))
|
PORT = int(os.getenv("PORT", "8080"))
|
||||||
STRATEGY_ID = os.getenv("STRATEGY_ID", "")
|
STRATEGY_ID = os.getenv("STRATEGY_ID", "")
|
||||||
|
XC_TOKEN = os.getenv("MODELHUB_XC_TOKEN", "")
|
||||||
|
TARGET_GPU = os.getenv("MODELHUB_TARGET_GPU", DEFAULT_TARGET_GPU)
|
||||||
|
AUTO_SCAN_ENABLED = os.getenv("AUTO_SCAN_ENABLED", "true").lower() not in {
|
||||||
|
"0", "false", "no", "off"
|
||||||
|
}
|
||||||
|
AUTO_SUBMIT_ENABLED = os.getenv("AUTO_SUBMIT_ENABLED", "true").lower() not in {
|
||||||
|
"0", "false", "no", "off"
|
||||||
|
}
|
||||||
|
SCAN_INTERVAL_SECONDS = max(int(os.getenv("SCAN_INTERVAL_SECONDS", "3600")), 60)
|
||||||
|
CANDIDATE_ENDPOINT = "https://modelhub.org.cn/api/computility/models/top/models"
|
||||||
|
TASK_PAGE_ENDPOINT = "https://modelhub.org.cn/api/adapt/task/page"
|
||||||
|
TASK_ADD_ENDPOINT = "https://modelhub.org.cn/api/adapt/task/add"
|
||||||
|
DEFAULT_CONFIG_PARAMS = "framework: vllm\napi: completion\nlang: zh\nmax_model_len: 2048\nmax_tokens: 256\ntemperature: 0.1\nrepetition_penalty: 1.0\ntop_p: 0.9\nsut_config:\n gpu_num: 1\n values:\n command:\n - python3\n - -m\n - vllm.entrypoints.openai.api_server\n - --host\n - 0.0.0.0\n - --port\n - '20644'\n - --served-model-name\n - llm\n - --model\n - /model\n - --max-model-len\n - '2048'\n - --tensor-parallel-size\n - '1'\n - --max-num-seqs\n - '8'\n - --enforce-eager\n - --disable-log-requests\n - --enable-prefix-caching\n - --trust-remote-code\nref_config:\n gpu_num: 1\n values:\n command:\n - vllm\n - serve\n - /model\n - --port\n - '80'\n - --served-model-name\n - llm\n - --max-model-len\n - '2048'\n - -tp\n - '1'\n - --enforce-eager\n - --trust-remote-code\n"
|
||||||
|
CONFIG_PARAMS = os.getenv("MODELHUB_CONFIG_PARAMS", DEFAULT_CONFIG_PARAMS)
|
||||||
MAX_BODY_BYTES = 1_000_000
|
MAX_BODY_BYTES = 1_000_000
|
||||||
|
SECRET_ASSIGNMENT_RE = re.compile(
|
||||||
|
r"(?i)(token|secret|password|api[_-]?key)=([^&\s]+)"
|
||||||
|
)
|
||||||
|
HTTP_OPENER = build_opener(ProxyHandler({}))
|
||||||
|
SCANNER_LOCK = threading.Lock()
|
||||||
|
SCANNER_STATE: dict[str, Any] = {
|
||||||
|
"enabled": AUTO_SCAN_ENABLED,
|
||||||
|
"running": False,
|
||||||
|
"last_scan_at": None,
|
||||||
|
"last_candidate": None,
|
||||||
|
"last_action": "not_started",
|
||||||
|
"last_error_type": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class ValidationError(ValueError):
|
class ValidationError(ValueError):
|
||||||
"""Raised when a request field is invalid."""
|
"""Raised when a request field is invalid."""
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_model_name(value: Any) -> str:
|
||||||
|
"""Return a bounded model identifier without URL credentials or query data."""
|
||||||
|
|
||||||
|
text = str(value or "未指定模型").strip()
|
||||||
|
if "://" in text:
|
||||||
|
try:
|
||||||
|
parsed = urlsplit(text)
|
||||||
|
text = f"{parsed.scheme}://{parsed.hostname or ''}{parsed.path}"
|
||||||
|
except ValueError:
|
||||||
|
text = "无效模型地址"
|
||||||
|
text = SECRET_ASSIGNMENT_RE.sub(r"\1=[REDACTED]", text)
|
||||||
|
return text[:200] or "未指定模型"
|
||||||
|
|
||||||
|
|
||||||
|
def _log_event(event: str, **fields: Any) -> None:
|
||||||
|
"""Write one structured, unbuffered log record to stdout."""
|
||||||
|
|
||||||
|
record = {
|
||||||
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||||
|
"event": event,
|
||||||
|
"agent": AGENT_NAME,
|
||||||
|
"version": AGENT_VERSION,
|
||||||
|
**fields,
|
||||||
|
}
|
||||||
|
print(json.dumps(record, ensure_ascii=False, separators=(",", ":")), flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _scanner_state_update(**fields: Any) -> None:
|
||||||
|
with SCANNER_LOCK:
|
||||||
|
SCANNER_STATE.update(fields)
|
||||||
|
|
||||||
|
|
||||||
|
def scanner_snapshot() -> dict[str, Any]:
|
||||||
|
with SCANNER_LOCK:
|
||||||
|
return dict(SCANNER_STATE)
|
||||||
|
|
||||||
|
|
||||||
|
def _http_json(
|
||||||
|
method: str,
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
payload: dict[str, Any] | None = None,
|
||||||
|
token: str = "",
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Call one ModelHub JSON endpoint without logging headers or response bodies."""
|
||||||
|
|
||||||
|
body = None
|
||||||
|
headers = {"Accept": "application/json"}
|
||||||
|
if payload is not None:
|
||||||
|
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||||||
|
headers["Content-Type"] = "application/json"
|
||||||
|
if token:
|
||||||
|
headers["Xc-Token"] = token
|
||||||
|
request = Request(url, data=body, headers=headers, method=method)
|
||||||
|
with HTTP_OPENER.open(request, timeout=15) as response:
|
||||||
|
result = json.loads(response.read().decode("utf-8"))
|
||||||
|
if not isinstance(result, dict):
|
||||||
|
raise RuntimeError("unexpected_response_shape")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _candidate_records(payload: dict[str, Any]) -> list[dict[str, str]]:
|
||||||
|
"""Extract text-generation candidates from the public ModelHub response."""
|
||||||
|
|
||||||
|
data = payload.get("data")
|
||||||
|
if not isinstance(data, list):
|
||||||
|
return []
|
||||||
|
candidates: list[dict[str, str]] = []
|
||||||
|
for item in data:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
continue
|
||||||
|
model_id = str(item.get("modelId") or "").strip()
|
||||||
|
info = item.get("taskLevelsInfo")
|
||||||
|
info = info if isinstance(info, dict) else {}
|
||||||
|
raw_type = str(
|
||||||
|
info.get("taskLevelName")
|
||||||
|
or info.get("taskLevelCode")
|
||||||
|
or info.get("taskLevelEnglishName")
|
||||||
|
or ""
|
||||||
|
).lower()
|
||||||
|
chinese_type = str(info.get("taskLevelChineseName") or "")
|
||||||
|
if not model_id or not (
|
||||||
|
"text-generation" in raw_type
|
||||||
|
or "text_generation" in raw_type
|
||||||
|
or chinese_type.strip() == "文本生成"
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
candidates.append({"model_id": model_id, "task_type": "text-generation"})
|
||||||
|
return candidates
|
||||||
|
|
||||||
|
|
||||||
|
def _model_size_score(model_id: str) -> float:
|
||||||
|
match = re.search(r"(?i)(\d+(?:\.\d+)?)\s*([bm])(?:\b|[_-])", model_id)
|
||||||
|
if not match:
|
||||||
|
return 1_000_000.0
|
||||||
|
size = float(match.group(1))
|
||||||
|
return size * (1_000 if match.group(2).lower() == "b" else 1)
|
||||||
|
|
||||||
|
|
||||||
|
def select_candidate(payload: dict[str, Any]) -> dict[str, str] | None:
|
||||||
|
"""Prefer a small official-EngineX text model from the public hot list."""
|
||||||
|
|
||||||
|
candidates = _candidate_records(payload)
|
||||||
|
if not candidates:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def score(item: dict[str, str]) -> tuple[int, float, str]:
|
||||||
|
model_id = item["model_id"].lower()
|
||||||
|
supported_family = any(
|
||||||
|
name in model_id for name in ("qwen3", "llama3", "deepseek-r1-distill")
|
||||||
|
)
|
||||||
|
return (0 if supported_family else 1, _model_size_score(model_id), model_id)
|
||||||
|
|
||||||
|
return min(candidates, key=score)
|
||||||
|
|
||||||
|
|
||||||
|
def _model_address(model_id: str) -> str:
|
||||||
|
safe_path = "/".join(quote(part, safe="") for part in model_id.split("/"))
|
||||||
|
return f"https://www.modelscope.cn/models/{safe_path}"
|
||||||
|
|
||||||
|
|
||||||
|
def scan_once(request_json: Any = None) -> dict[str, Any]:
|
||||||
|
"""Discover one candidate, deduplicate it, and optionally submit one task."""
|
||||||
|
|
||||||
|
call = request_json or _http_json
|
||||||
|
scan_at = datetime.now(timezone.utc).isoformat()
|
||||||
|
_scanner_state_update(running=True, last_scan_at=scan_at, last_error_type=None)
|
||||||
|
_log_event("candidate_scan_started", source="modelhub_hot_models")
|
||||||
|
candidate_payload = call("GET", CANDIDATE_ENDPOINT)
|
||||||
|
candidate = select_candidate(candidate_payload)
|
||||||
|
if candidate is None:
|
||||||
|
result = {"action": "no_candidate"}
|
||||||
|
_scanner_state_update(running=False, last_candidate=None, last_action=result["action"])
|
||||||
|
_log_event("candidate_scan_completed", candidate_count=0, action=result["action"])
|
||||||
|
return result
|
||||||
|
|
||||||
|
model_id = _safe_model_name(candidate["model_id"])
|
||||||
|
_scanner_state_update(last_candidate=model_id)
|
||||||
|
_log_event(
|
||||||
|
"candidate_discovered",
|
||||||
|
model=model_id,
|
||||||
|
task_type=candidate["task_type"],
|
||||||
|
target_gpu_configured=bool(TARGET_GPU),
|
||||||
|
)
|
||||||
|
|
||||||
|
missing = []
|
||||||
|
if not AUTO_SUBMIT_ENABLED:
|
||||||
|
missing.append("auto_submit_disabled")
|
||||||
|
if not XC_TOKEN:
|
||||||
|
missing.append("xc_token_missing")
|
||||||
|
if not STRATEGY_ID:
|
||||||
|
missing.append("strategy_id_missing")
|
||||||
|
if not TARGET_GPU:
|
||||||
|
missing.append("target_gpu_missing")
|
||||||
|
if not CONFIG_PARAMS:
|
||||||
|
missing.append("task_config_missing")
|
||||||
|
if missing:
|
||||||
|
result = {"action": "submission_skipped", "reasons": missing, "model": model_id}
|
||||||
|
_scanner_state_update(running=False, last_action=result["action"])
|
||||||
|
_log_event("task_submission_skipped", model=model_id, reasons=missing)
|
||||||
|
return result
|
||||||
|
|
||||||
|
query = urlencode(
|
||||||
|
{
|
||||||
|
"current": 1,
|
||||||
|
"pageSize": 1,
|
||||||
|
"onlyMine": "true",
|
||||||
|
"modelId": candidate["model_id"],
|
||||||
|
"gpuType": TARGET_GPU,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
history = call("GET", f"{TASK_PAGE_ENDPOINT}?{query}", token=XC_TOKEN)
|
||||||
|
history_data = history.get("data") if isinstance(history, dict) else None
|
||||||
|
records = history_data.get("records") if isinstance(history_data, dict) else []
|
||||||
|
if records:
|
||||||
|
result = {"action": "duplicate_skipped", "model": model_id}
|
||||||
|
_scanner_state_update(running=False, last_action=result["action"])
|
||||||
|
_log_event("task_duplicate_skipped", model=model_id, target_gpu=TARGET_GPU)
|
||||||
|
return result
|
||||||
|
|
||||||
|
task_payload = {
|
||||||
|
"modelAddress": _model_address(candidate["model_id"]),
|
||||||
|
"taskType": candidate["task_type"],
|
||||||
|
"targetGpu": TARGET_GPU,
|
||||||
|
"framework": "vllm",
|
||||||
|
"strategyId": STRATEGY_ID,
|
||||||
|
"configParams": CONFIG_PARAMS,
|
||||||
|
}
|
||||||
|
response = call("POST", TASK_ADD_ENDPOINT, payload=task_payload, token=XC_TOKEN)
|
||||||
|
response_data = response.get("data") if isinstance(response, dict) else None
|
||||||
|
task_id = response_data.get("taskId") if isinstance(response_data, dict) else response_data
|
||||||
|
result = {"action": "submitted", "model": model_id, "task_id": task_id}
|
||||||
|
_scanner_state_update(running=False, last_action=result["action"])
|
||||||
|
_log_event(
|
||||||
|
"task_submitted",
|
||||||
|
model=model_id,
|
||||||
|
target_gpu=TARGET_GPU,
|
||||||
|
task_id=task_id,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _scanner_loop() -> None:
|
||||||
|
_log_event(
|
||||||
|
"scanner_started",
|
||||||
|
interval_seconds=SCAN_INTERVAL_SECONDS,
|
||||||
|
auto_submit=AUTO_SUBMIT_ENABLED,
|
||||||
|
token_present=bool(XC_TOKEN),
|
||||||
|
target_gpu_configured=bool(TARGET_GPU),
|
||||||
|
)
|
||||||
|
while not STOP.is_set():
|
||||||
|
try:
|
||||||
|
scan_once()
|
||||||
|
except Exception as exc:
|
||||||
|
error_type = type(exc).__name__
|
||||||
|
_scanner_state_update(
|
||||||
|
running=False,
|
||||||
|
last_action="scan_error",
|
||||||
|
last_error_type=error_type,
|
||||||
|
)
|
||||||
|
_log_event("scanner_error", error_type=error_type)
|
||||||
|
if STOP.wait(SCAN_INTERVAL_SECONDS):
|
||||||
|
break
|
||||||
|
|
||||||
|
|
||||||
def _first(payload: dict[str, Any], *names: str) -> Any:
|
def _first(payload: dict[str, Any], *names: str) -> Any:
|
||||||
for name in names:
|
for name in names:
|
||||||
if name in payload and payload[name] not in (None, ""):
|
if name in payload and payload[name] not in (None, ""):
|
||||||
@@ -160,7 +417,7 @@ def analyze(payload: dict[str, Any]) -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
class Handler(BaseHTTPRequestHandler):
|
class Handler(BaseHTTPRequestHandler):
|
||||||
server_version = "ModelHubCardAdvisor/1.0"
|
server_version = "ModelHubCardAdvisor/1.2"
|
||||||
|
|
||||||
def _json(self, payload: dict[str, Any], status: int = 200) -> None:
|
def _json(self, payload: dict[str, Any], status: int = 200) -> None:
|
||||||
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||||||
@@ -172,7 +429,17 @@ class Handler(BaseHTTPRequestHandler):
|
|||||||
|
|
||||||
def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract
|
def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract
|
||||||
if self.path == "/health":
|
if self.path == "/health":
|
||||||
self._json({"status": "ok", "agent": AGENT_NAME, "version": AGENT_VERSION})
|
self._json(
|
||||||
|
{
|
||||||
|
"status": "ok",
|
||||||
|
"agent": AGENT_NAME,
|
||||||
|
"version": AGENT_VERSION,
|
||||||
|
"scanner": scanner_snapshot(),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if self.path == "/scanner":
|
||||||
|
self._json(scanner_snapshot())
|
||||||
return
|
return
|
||||||
if self.path == "/":
|
if self.path == "/":
|
||||||
self._json(
|
self._json(
|
||||||
@@ -182,8 +449,20 @@ class Handler(BaseHTTPRequestHandler):
|
|||||||
"description": f"{CARD_VENDOR} {CARD_MODEL} 专属部署前预检",
|
"description": f"{CARD_VENDOR} {CARD_MODEL} 专属部署前预检",
|
||||||
"official_target": {"vendor": CARD_VENDOR, "model": CARD_MODEL},
|
"official_target": {"vendor": CARD_VENDOR, "model": CARD_MODEL},
|
||||||
"strategy_id_present": bool(STRATEGY_ID),
|
"strategy_id_present": bool(STRATEGY_ID),
|
||||||
"endpoints": ["GET /health", "POST /analyze", "POST /task"],
|
"target_gpu": TARGET_GPU or None,
|
||||||
"external_writes": False,
|
"endpoints": [
|
||||||
|
"GET /health",
|
||||||
|
"GET /scanner",
|
||||||
|
"POST /analyze",
|
||||||
|
"POST /task",
|
||||||
|
],
|
||||||
|
"external_writes_enabled": bool(
|
||||||
|
AUTO_SUBMIT_ENABLED
|
||||||
|
and XC_TOKEN
|
||||||
|
and STRATEGY_ID
|
||||||
|
and TARGET_GPU
|
||||||
|
and CONFIG_PARAMS
|
||||||
|
),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -198,24 +477,52 @@ class Handler(BaseHTTPRequestHandler):
|
|||||||
if length <= 0:
|
if length <= 0:
|
||||||
raise ValidationError("请求正文不能为空")
|
raise ValidationError("请求正文不能为空")
|
||||||
if length > MAX_BODY_BYTES:
|
if length > MAX_BODY_BYTES:
|
||||||
|
_log_event("request_rejected", path=self.path, reason="payload_too_large")
|
||||||
self._json({"error": "payload_too_large"}, 413)
|
self._json({"error": "payload_too_large"}, 413)
|
||||||
return
|
return
|
||||||
payload = json.loads(self.rfile.read(length))
|
payload = json.loads(self.rfile.read(length))
|
||||||
self._json(analyze(payload))
|
model_name = _safe_model_name(
|
||||||
|
_first(payload, "model_name", "model", "model_address")
|
||||||
|
if isinstance(payload, dict)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
_log_event("analysis_started", path=self.path, model=model_name)
|
||||||
|
result = analyze(payload)
|
||||||
|
_log_event(
|
||||||
|
"analysis_completed",
|
||||||
|
path=self.path,
|
||||||
|
model=model_name,
|
||||||
|
verdict=result["verdict"],
|
||||||
|
target_matches=result["target_matches"],
|
||||||
|
)
|
||||||
|
self._json(result)
|
||||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||||
|
_log_event("request_rejected", path=self.path, reason="invalid_json")
|
||||||
self._json({"error": "invalid_json", "message": "请求正文必须是有效 JSON"}, 400)
|
self._json({"error": "invalid_json", "message": "请求正文必须是有效 JSON"}, 400)
|
||||||
except ValidationError as exc:
|
except ValidationError as exc:
|
||||||
|
_log_event(
|
||||||
|
"request_rejected",
|
||||||
|
path=self.path,
|
||||||
|
reason="validation_error",
|
||||||
|
message=str(exc),
|
||||||
|
)
|
||||||
self._json({"error": "validation_error", "message": str(exc)}, 422)
|
self._json({"error": "validation_error", "message": str(exc)}, 422)
|
||||||
|
|
||||||
def log_message(self, _format: str, *_args: Any) -> None:
|
def log_message(self, _format: str, *args: Any) -> None:
|
||||||
return
|
status = str(args[1]) if len(args) > 1 else "unknown"
|
||||||
|
_log_event(
|
||||||
|
"http_access",
|
||||||
|
method=self.command,
|
||||||
|
path=urlsplit(self.path).path,
|
||||||
|
status=status,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
STOP = threading.Event()
|
STOP = threading.Event()
|
||||||
|
|
||||||
|
|
||||||
def _handle_signal(signum: int, _frame: Any) -> None:
|
def _handle_signal(signum: int, _frame: Any) -> None:
|
||||||
print(f"received signal {signum}; shutting down", flush=True)
|
_log_event("shutdown_requested", signal=signum)
|
||||||
STOP.set()
|
STOP.set()
|
||||||
|
|
||||||
|
|
||||||
@@ -224,10 +531,29 @@ def main() -> None:
|
|||||||
signal.signal(signal.SIGINT, _handle_signal)
|
signal.signal(signal.SIGINT, _handle_signal)
|
||||||
server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
|
server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
|
||||||
server.timeout = 1
|
server.timeout = 1
|
||||||
print(f"{AGENT_NAME} {AGENT_VERSION} listening on 0.0.0.0:{PORT}", flush=True)
|
_log_event(
|
||||||
|
"service_started",
|
||||||
|
host="0.0.0.0",
|
||||||
|
port=PORT,
|
||||||
|
target_vendor=CARD_VENDOR,
|
||||||
|
target_model=CARD_MODEL,
|
||||||
|
)
|
||||||
|
scanner_thread = None
|
||||||
|
if AUTO_SCAN_ENABLED:
|
||||||
|
scanner_thread = threading.Thread(
|
||||||
|
target=_scanner_loop,
|
||||||
|
name="modelhub-candidate-scanner",
|
||||||
|
daemon=True,
|
||||||
|
)
|
||||||
|
scanner_thread.start()
|
||||||
|
else:
|
||||||
|
_scanner_state_update(last_action="disabled")
|
||||||
|
_log_event("scanner_disabled", reason="auto_scan_disabled")
|
||||||
while not STOP.is_set():
|
while not STOP.is_set():
|
||||||
server.handle_request()
|
server.handle_request()
|
||||||
server.server_close()
|
server.server_close()
|
||||||
|
if scanner_thread is not None:
|
||||||
|
scanner_thread.join(timeout=5)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
115
test_main.py
115
test_main.py
@@ -1,6 +1,22 @@
|
|||||||
|
import io
|
||||||
|
import json
|
||||||
import unittest
|
import unittest
|
||||||
|
from contextlib import redirect_stdout
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
from main import CARD_MODEL, CARD_VENDOR, ValidationError, analyze
|
import main as agent_main
|
||||||
|
|
||||||
|
from main import (
|
||||||
|
AGENT_NAME,
|
||||||
|
CARD_MODEL,
|
||||||
|
CARD_VENDOR,
|
||||||
|
ValidationError,
|
||||||
|
_log_event,
|
||||||
|
_safe_model_name,
|
||||||
|
analyze,
|
||||||
|
scan_once,
|
||||||
|
select_candidate,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class AnalyzeTests(unittest.TestCase):
|
class AnalyzeTests(unittest.TestCase):
|
||||||
@@ -45,6 +61,103 @@ class AnalyzeTests(unittest.TestCase):
|
|||||||
with self.assertRaises(ValidationError):
|
with self.assertRaises(ValidationError):
|
||||||
analyze({"cards": 0})
|
analyze({"cards": 0})
|
||||||
|
|
||||||
|
def test_model_log_identifier_removes_url_credentials(self):
|
||||||
|
value = "https://user:pass@example.com/org/model?token=secret#fragment"
|
||||||
|
self.assertEqual(_safe_model_name(value), "https://example.com/org/model")
|
||||||
|
|
||||||
|
def test_structured_log_is_valid_json(self):
|
||||||
|
stream = io.StringIO()
|
||||||
|
with redirect_stdout(stream):
|
||||||
|
_log_event("analysis_completed", model="example/model", verdict="ready")
|
||||||
|
record = json.loads(stream.getvalue())
|
||||||
|
self.assertEqual(record["agent"], AGENT_NAME)
|
||||||
|
self.assertEqual(record["event"], "analysis_completed")
|
||||||
|
self.assertEqual(record["model"], "example/model")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def candidate_payload():
|
||||||
|
return {
|
||||||
|
"code": 0,
|
||||||
|
"data": [
|
||||||
|
{
|
||||||
|
"modelId": "Qwen/Qwen3-30B-Instruct",
|
||||||
|
"taskLevelsInfo": {"taskLevelChineseName": "文本生成"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"modelId": "Qwen/Qwen3-4B-Instruct-2507",
|
||||||
|
"taskLevelsInfo": {"taskLevelChineseName": "文本生成"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"modelId": "example/image-model",
|
||||||
|
"taskLevelsInfo": {"taskLevelChineseName": "文本生成图片"},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_candidate_selection_prefers_smaller_supported_text_model(self):
|
||||||
|
candidate = select_candidate(self.candidate_payload())
|
||||||
|
self.assertIsNotNone(candidate)
|
||||||
|
self.assertEqual(candidate["model_id"], "Qwen/Qwen3-4B-Instruct-2507")
|
||||||
|
|
||||||
|
def test_scan_skips_write_without_private_token(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_request(method, url, **kwargs):
|
||||||
|
calls.append((method, url, kwargs))
|
||||||
|
return self.candidate_payload()
|
||||||
|
|
||||||
|
with patch.object(agent_main, "XC_TOKEN", ""):
|
||||||
|
result = scan_once(fake_request)
|
||||||
|
self.assertEqual(result["action"], "submission_skipped")
|
||||||
|
self.assertIn("xc_token_missing", result["reasons"])
|
||||||
|
self.assertEqual([call[0] for call in calls], ["GET"])
|
||||||
|
|
||||||
|
def test_scan_deduplicates_before_submission(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_request(method, url, **kwargs):
|
||||||
|
calls.append((method, url, kwargs))
|
||||||
|
if "top/models" in url:
|
||||||
|
return self.candidate_payload()
|
||||||
|
return {"code": 0, "data": {"records": [{"taskId": "existing"}]}}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(agent_main, "XC_TOKEN", "private-token"),
|
||||||
|
patch.object(agent_main, "STRATEGY_ID", "strategy-id"),
|
||||||
|
patch.object(agent_main, "TARGET_GPU", "verified-gpu"),
|
||||||
|
patch.object(agent_main, "CONFIG_PARAMS", "framework: vllm"),
|
||||||
|
patch.object(agent_main, "AUTO_SUBMIT_ENABLED", True),
|
||||||
|
):
|
||||||
|
result = scan_once(fake_request)
|
||||||
|
self.assertEqual(result["action"], "duplicate_skipped")
|
||||||
|
self.assertEqual([call[0] for call in calls], ["GET", "GET"])
|
||||||
|
|
||||||
|
def test_scan_submits_exactly_one_task_after_empty_history(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_request(method, url, **kwargs):
|
||||||
|
calls.append((method, url, kwargs))
|
||||||
|
if "top/models" in url:
|
||||||
|
return self.candidate_payload()
|
||||||
|
if method == "GET":
|
||||||
|
return {"code": 0, "data": {"records": []}}
|
||||||
|
return {"code": 0, "data": {"taskId": "new-task"}}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(agent_main, "XC_TOKEN", "private-token"),
|
||||||
|
patch.object(agent_main, "STRATEGY_ID", "strategy-id"),
|
||||||
|
patch.object(agent_main, "TARGET_GPU", "verified-gpu"),
|
||||||
|
patch.object(agent_main, "CONFIG_PARAMS", "framework: vllm"),
|
||||||
|
patch.object(agent_main, "AUTO_SUBMIT_ENABLED", True),
|
||||||
|
):
|
||||||
|
result = scan_once(fake_request)
|
||||||
|
self.assertEqual(result["action"], "submitted")
|
||||||
|
self.assertEqual(result["task_id"], "new-task")
|
||||||
|
self.assertEqual([call[0] for call in calls], ["GET", "GET", "POST"])
|
||||||
|
submitted = calls[-1][2]["payload"]
|
||||||
|
self.assertEqual(submitted["targetGpu"], "verified-gpu")
|
||||||
|
self.assertEqual(submitted["strategyId"], "strategy-id")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user