feat: 模型推理异步加载与对话链路修复,同步基线

模型推理全异步化改造:
- 计算节点 InferenceSession 改为后台线程异步加载模型,load 立即返回,
  加载期间事件循环保持响应(/inference/status 与 /health 不阻塞)
- 后端模型加载改为异步派发 + 轮询对账器(reconcile_inference_loads),
  任务状态由 starting 自动推进到 ready/error,解决多节点启动超时
  (timeout of 120000ms exceeded)
- 推理删除/卸载改为任务感知 + 短超时,删除先删记录再 best-effort 卸载,
  不再被不可达节点阻塞;同节点新模型替换旧任务标记失效
- 流式对话透传 task_id/node_id 路由到真正加载模型的算力节点,
  useStreamChat 解析 SSE 错误帧以干净文案展示
- 对话历史按任务 id 本地持久化,退出重进可恢复;移除页脚提示文本
- 新增后端推理异步加载与计算节点异步状态机单元测试

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
wuyongtao
2026-08-04 16:59:34 +08:00
parent 250e060271
commit 0271942ba5
21 changed files with 1272 additions and 245 deletions

View File

@@ -33,6 +33,15 @@ def _select_first_online_node(store: Any) -> dict[str, Any] | None:
return None
def _candidate_online_nodes(store: Any, preferred_node_id: str | None = None) -> list[dict[str, Any]]:
nodes = [node for node in store.compute_nodes() if node.get("enabled") and node.get("scheduler_status") == "online"]
if not preferred_node_id:
return nodes
preferred = [node for node in nodes if node.get("id") == preferred_node_id]
others = [node for node in nodes if node.get("id") != preferred_node_id]
return preferred + others
def _build_messages_payload(payload: dict[str, Any]) -> dict[str, Any]:
"""Convert frontend inference payload to compute API messages format.
@@ -61,10 +70,47 @@ def _build_messages_payload(payload: dict[str, Any]) -> dict[str, Any]:
}
def _node_for_inference_payload(store: Any, payload: dict[str, Any]) -> dict[str, Any] | None:
node_id = payload.get("node_id") or payload.get("compute_node_id")
task_id = payload.get("task_id") or payload.get("compare_task_id")
if task_id and not node_id:
try:
task = store.compare_task(str(task_id))
load_status = task.get("load_status") or {}
if isinstance(load_status, str):
load_status = json.loads(load_status)
loaded_models = load_status.get("loaded_models") or []
ready_model = next((item for item in loaded_models if item.get("status") in {"ready", "running"} and item.get("node_id")), None)
if ready_model:
node_id = ready_model.get("node_id")
except Exception:
node_id = None
if node_id:
return next((node for node in store.compute_nodes() if node.get("id") == node_id), None)
return _select_first_online_node(store)
async def _stream_chat_proxy(payload: dict[str, Any]) -> StreamingResponse:
"""Common SSE streaming proxy: convert payload → forward to compute node → stream back."""
store = get_platform_store()
node = _select_first_online_node(store)
# 任务仍在加载中时,直接返回明确的加载中提示,避免转发到尚未就绪的节点
task_id = payload.get("task_id") or payload.get("compare_task_id")
if task_id:
try:
task = store.compare_task(str(task_id))
load_status = task.get("load_status") or {}
if isinstance(load_status, str):
load_status = json.loads(load_status)
items = load_status.get("loaded_models") or []
if items and not any(item.get("status") in {"ready", "running"} for item in items):
if any(item.get("status") == "starting" for item in items):
return StreamingResponse(
iter(['data: {"error": "模型加载中,请稍候再试"}\n\n']),
media_type="text/event-stream",
)
except Exception: # noqa: BLE001 - fall through to normal routing on lookup errors
pass
node = _node_for_inference_payload(store, payload)
if not node:
return StreamingResponse(
iter(['data: {"error": "no online compute node available for inference"}\n\n']),
@@ -729,12 +775,17 @@ async def merge_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
None,
)
base_model_path = payload.get("base_model_path") or (trained_model and trained_model.get("base_model_path"))
adapter_path = payload.get("adapter_path") or payload.get("adapter_name_or_path") or (trained_model and trained_model.get("merged_path"))
adapter_path = (
payload.get("adapter_path")
or payload.get("adapter_name_or_path")
or (trained_model and (trained_model.get("artifact_dir") or trained_model.get("adapter_path") or trained_model.get("merged_path")))
)
if not base_model_path:
raise fail(400, "base_model_path is required")
if not adapter_path:
raise fail(400, "adapter_path is required")
node = store.schedule_node({**payload, "gpus": payload.get("gpus") or []})
requested_node_id = payload.get("requested_node_id") or payload.get("compute_node_id") or (trained_model and trained_model.get("compute_node_id"))
node = store.schedule_node({**payload, "requested_node_id": requested_node_id, "gpus": payload.get("gpus") or []})
health = node.get("health_detail") or {}
output_root = str(health.get("output_root") or f"{node['data_root'].rstrip('/')}/outputs")
output_name = str(payload.get("output_model_name") or payload.get("merged_model_name") or f"{trained_model_id or 'model'}-merged")
@@ -753,11 +804,13 @@ async def merge_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
"gpus": payload.get("gpus") or [],
"trained_model_id": trained_model["id"] if trained_model else trained_model_id,
"model_name": trained_model["name"] if trained_model else payload.get("model_name"),
"compute_node_id": node["id"],
"compute_node_code": node.get("code"),
}
if get_settings().compute_mode == "simulator":
job = {"id": job_payload["id"], "status": "queued", "progress": 10, "command": [], "output_dir": output_dir}
else:
client = ComputeNodeClient(node["api_base_url"])
client = ComputeNodeClient(node["api_base_url"], timeout=900)
preview = await client.validate_job(job_payload)
if not preview.get("valid", False):
raise fail(409, "; ".join(preview.get("errors") or ["merge preflight failed"]))
@@ -1537,28 +1590,49 @@ async def model_compare_detail(task_id: str) -> dict[str, Any]:
raise fail(404, "compare task not found")
async def _unload_from_compute_node() -> dict[str, Any]:
"""Best-effort unload the inference model from the first online compute node."""
store = get_platform_store()
# Clear all inference tracking — only one model can be loaded at a time
for node in store.compute_nodes():
store.mark_inference_unloaded(node["id"])
node = _select_first_online_node(store)
if not node:
return {"unloaded": False, "error": "no online compute node"}
try:
client = ComputeNodeClient(node["api_base_url"])
result = await client._request("POST", "/inference/unload", json_data={})
return result
except Exception as exc:
return {"unloaded": False, "error": str(exc)}
async def _unload_from_compute_node(store: Any, task: dict[str, Any] | None = None) -> dict[str, Any]:
"""Best-effort unload the inference model from the node(s) that hold it.
任务感知:优先卸载 ``task.load_status.loaded_models`` 中记录的节点;
无任务时回退到平台记录的已加载推理的节点。每个节点使用短超时,
保证卸载永远不会长时间阻塞调用方(例如删除操作)。
"""
node_ids: set[str] = set()
if task:
load_status = task.get("load_status") or {}
if isinstance(load_status, str):
try:
load_status = json.loads(load_status)
except json.JSONDecodeError:
load_status = {}
node_ids = {item.get("node_id") for item in load_status.get("loaded_models") or [] if item.get("node_id")}
if not node_ids:
node_ids = {node["id"] for node in store.compute_nodes() if store.is_inference_loaded(node["id"])}
nodes = [node for node in store.compute_nodes() if node["id"] in node_ids]
results: list[dict[str, Any]] = []
for node in nodes:
try:
result = await ComputeNodeClient(node["api_base_url"]).inference_unload()
results.append({"node_id": node["id"], "node_code": node.get("code"), "success": True, "result": result})
except Exception as exc: # noqa: BLE001 - best-effort unload must not raise
results.append({"node_id": node["id"], "node_code": node.get("code"), "success": False, "error": str(exc)})
finally:
store.mark_inference_unloaded(node["id"])
return {"unloaded": bool(results), "nodes": results}
@router.delete("/model-compare/{task_id}")
async def model_compare_delete(task_id: str) -> dict[str, Any]:
# 删除前先释放算力节点上的模型
await _unload_from_compute_node()
# 先删记录(快),再 best-effort 释放算力节点上的模型——删除绝不被卸载阻塞
try:
task = get_platform_store().compare_task(task_id)
except KeyError:
raise fail(404, "compare task not found")
get_platform_store().delete_compare_task(task_id)
try:
await _unload_from_compute_node(get_platform_store(), task=task)
except Exception: # noqa: BLE001 - deletion must succeed even if unload fails
pass
return ok({"deleted": task_id})
@@ -1585,9 +1659,44 @@ async def model_compare_update_load_status(task_id: str, payload: dict[str, Any]
raise fail(404, "compare task not found")
def _invalidate_superseded_models(store: Any, task_id: str, loaded_models: list[dict[str, Any]]) -> None:
"""同一计算节点同一时刻只能加载一个推理模型。
当新任务把模型派发到了某节点后,把其它任务中在该节点上 ready/running
的模型标记为已被替换,保持平台 DB 与计算节点实际状态一致。
"""
taken_node_ids = {m.get("node_id") for m in loaded_models if m.get("node_id") and m.get("status") == "starting"}
if not taken_node_ids:
return
for other in store.compare_tasks():
if str(other.get("id")) == str(task_id):
continue
load_status = other.get("load_status") or {}
if isinstance(load_status, str):
try:
load_status = json.loads(load_status)
except json.JSONDecodeError:
load_status = {}
items = load_status.get("loaded_models") or []
changed = False
for item in items:
if item.get("node_id") in taken_node_ids and item.get("status") in {"ready", "running"}:
item["status"] = "error"
item["error"] = "模型已被其他推理任务替换"
changed = True
if changed:
new_status = "loaded" if any(i.get("status") in {"ready", "running"} for i in items) else "failed"
store.update_compare_task(other["id"], {"status": new_status, "load_status": {"loaded_models": items}})
@router.post("/model-compare/{task_id}/load")
async def model_compare_load(task_id: str) -> dict[str, Any]:
"""真正加载模型到算力节点(不再使用假 PID/端口)。"""
"""异步派发模型加载到算力节点,立即返回。
加载进度由轮询对账器compute_poller → reconcile_inference_loads推进
任务项先以 status=starting 记录,对账器查询节点 /inference/status 后
推进到 ready/error。这里只负责把加载请求派发出去绝不同步等待加载完成。
"""
try:
store = get_platform_store()
task = store.compare_task(task_id)
@@ -1597,15 +1706,14 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
models = json.loads(models)
except json.JSONDecodeError:
models = []
# 选取在线算力节点
node = _select_first_online_node(store)
if not node:
online_nodes = _candidate_online_nodes(store)
if not online_nodes:
return ok({"status": "failed", "error": "no online compute node"})
client = ComputeNodeClient(node["api_base_url"])
loaded_models = []
for item in models:
if not isinstance(item, dict):
continue
preferred_node_id = item.get("node_id") or item.get("compute_node_id")
model_path = item.get("model_path", "")
if not model_path:
# 尝试从模型库获取路径
@@ -1614,28 +1722,46 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
db_model = store.model(model_id)
model_path = db_model.get("path", "")
except KeyError:
pass
trained_model = next((m for m in store.trained_models() if str(m.get("id")) == str(model_id)), None)
if trained_model:
model_path = trained_model.get("merged_path") or trained_model.get("artifact_dir") or ""
preferred_node_id = preferred_node_id or trained_model.get("compute_node_id")
if not model_path:
loaded_models.append({**item, "status": "error", "error": "model_path not found"})
continue
# 真正调用算力节点加载模型
load_payload = {
"model_name_or_path": model_path,
"template": item.get("template", "qwen"),
}
if item.get("adapter_path"):
load_payload["adapter_name_or_path"] = item["adapter_path"]
try:
result = await client._request("POST", "/inference/load", json_data=load_payload)
if result.get("loaded"):
if get_settings().compute_mode == "simulator":
loaded_models.append({**item, "status": "ready", "node_id": "", "node_name": ""})
continue
# 只派发HTTP 响应成功即视为已接受节点会异步加载loaded 字段忽略
item_dispatched = False
errors = []
for node in _candidate_online_nodes(store, preferred_node_id):
try:
client = ComputeNodeClient(node["api_base_url"])
await client.inference_load(load_payload)
store.mark_inference_loaded(node["id"])
loaded_models.append({**item, "status": "ready"})
else:
loaded_models.append({**item, "status": "error", "error": result.get("error", "load failed")})
except Exception as exc:
loaded_models.append({**item, "status": "error", "error": str(exc)})
status = "loaded" if any(m.get("status") == "ready" for m in loaded_models) else "failed"
loaded_models.append({**item, "status": "starting", "node_id": node["id"], "node_name": node.get("name")})
item_dispatched = True
break
except Exception as exc: # noqa: BLE001 - try next candidate node
errors.append(f"{node.get('name') or node.get('code')}: {exc}")
if not item_dispatched:
loaded_models.append({**item, "status": "error", "error": "; ".join(errors) or "load dispatch failed"})
if any(m.get("status") == "starting" for m in loaded_models):
status = "starting"
elif any(m.get("status") == "error" for m in loaded_models):
status = "failed"
else:
status = "loaded"
updated = store.update_compare_task(task_id, {"status": status, "load_status": {"loaded_models": loaded_models}})
# 同一节点同一时刻只能有一个推理模型;新任务占用了节点后,把其它任务上该节点的模型标记为已被替换
_invalidate_superseded_models(store, task_id, loaded_models)
return ok(updated)
except KeyError:
raise fail(404, "compare task not found")
@@ -1645,8 +1771,9 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
async def model_compare_unload(task_id: str) -> dict[str, Any]:
try:
store = get_platform_store()
# 真正释放算力节点上的模型资源
unload_result = await _unload_from_compute_node()
task = store.compare_task(task_id)
# 任务感知卸载:只释放该任务实际加载到的节点,短超时快速返回
unload_result = await _unload_from_compute_node(store, task=task)
updated = store.update_compare_task(task_id, {"status": "pending", "load_status": {"loaded_models": []}})
return ok({"task": updated, "unload": unload_result})
except KeyError:
@@ -1662,7 +1789,7 @@ async def model_compare_start_model(task_id: str, payload: dict[str, Any] = Body
async def model_compare_chat_with_port(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
"""Proxy non-streaming chat to the compute node running the inference model."""
store = get_platform_store()
node = _select_first_online_node(store)
node = _node_for_inference_payload(store, payload)
if not node:
return ok({"response": "no online compute node available for inference", "request": payload})
try:
@@ -1717,8 +1844,9 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[
return ok({"loaded": False, "error": "no online compute node"})
try:
client = ComputeNodeClient(node["api_base_url"])
result = await client._request("POST", "/inference/load", json_data=payload)
if result.get("loaded"):
# 计算节点现在异步加载HTTP 接受loading/ready即视为派发成功
result = await client.inference_load(payload)
if result.get("loaded") or result.get("status") in {"loading", "ready"}:
store.mark_inference_loaded(node["id"])
return ok(result)
except Exception as exc:
@@ -1729,18 +1857,19 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[
async def model_chat_local_unload() -> dict[str, Any]:
"""Unload the inference model from the compute node."""
store = get_platform_store()
# Clear all inference tracking
# 释放所有已加载推理的节点短超时best-effort
results: list[dict[str, Any]] = []
for n in store.compute_nodes():
store.mark_inference_unloaded(n["id"])
node = _select_first_online_node(store)
if not node:
return ok({"unloaded": False, "error": "no online compute node"})
try:
client = ComputeNodeClient(node["api_base_url"])
result = await client._request("POST", "/inference/unload", json_data={})
return ok(result)
except Exception as exc:
return ok({"unloaded": False, "error": str(exc)})
if not store.is_inference_loaded(n["id"]):
continue
try:
result = await ComputeNodeClient(n["api_base_url"]).inference_unload()
results.append({"node_id": n["id"], "success": True, "result": result})
except Exception as exc: # noqa: BLE001 - best-effort unload
results.append({"node_id": n["id"], "success": False, "error": str(exc)})
finally:
store.mark_inference_unloaded(n["id"])
return ok({"unloaded": True, "nodes": results})
@router.get("/model-chat/local/status")
@@ -1752,7 +1881,7 @@ async def model_chat_local_status() -> dict[str, Any]:
return ok({"loaded": False, "error": "no online compute node"})
try:
client = ComputeNodeClient(node["api_base_url"])
result = await client._request("GET", "/inference/status")
result = await client.inference_status()
return ok(result)
except Exception as exc:
return ok({"loaded": False, "error": str(exc)})
@@ -1770,8 +1899,9 @@ async def model_chat_trained_preload(payload: dict[str, Any] = Body(...)) -> dic
return ok({"loaded": False, "error": "no online compute node"})
try:
client = ComputeNodeClient(node["api_base_url"])
result = await client._request("POST", "/inference/load", json_data=payload)
if result.get("loaded"):
# 计算节点现在异步加载HTTP 接受loading/ready即视为派发成功
result = await client.inference_load(payload)
if result.get("loaded") or result.get("status") in {"loading", "ready"}:
store.mark_inference_loaded(node["id"])
return ok(result)
except Exception as exc: