model_train_dm/backend/apps/training/consumers.py
2026-07-27 17:51:49 +08:00

151 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import asyncio
import json
import os
from channels.db import database_sync_to_async
from channels.generic.websocket import AsyncWebsocketConsumer
from django.conf import settings
from apps.core.models import AiAlgorithm, AiAlgorithmClass, AiAlgorithmTrainRecords
class TrainLogConsumer(AsyncWebsocketConsumer):
async def connect(self):
self.train_id = self.scope["url_route"]["kwargs"]["train_id"]
self.group_name = f"train_{self.train_id}"
await self.channel_layer.group_add(self.group_name, self.channel_name)
print(f"客户端[{self.train_id}],websocket连接服务器成功")
await self.accept()
self._stop_event = asyncio.Event()
self._tail_task = asyncio.create_task(self._tail_train_log())
async def disconnect(self, close_code):
if hasattr(self, "_stop_event"):
self._stop_event.set()
if hasattr(self, "_tail_task"):
self._tail_task.cancel()
try:
await self._tail_task
except Exception:
pass
await self.channel_layer.group_discard(self.group_name, self.channel_name)
async def receive(self, text_data):
pass
async def send_log(self, event):
message = event["message"]
await self.send(text_data=message)
@staticmethod
def _is_valid_json(text: str) -> bool:
try:
json.loads(text)
return True
except Exception:
return False
async def _send_json(self, payload):
await self.send(text_data=json.dumps(payload, ensure_ascii=False))
@database_sync_to_async
def _resolve_log_file_path(self) -> str:
ai_train = AiAlgorithmTrainRecords.objects.get(id=self.train_id)
ai_algorithm = AiAlgorithm.objects.get(id=ai_train.algorithm_id)
ai_algorithm_class = AiAlgorithmClass.objects.get(id=ai_algorithm.algorithm_class)
run_name = "train"
flat_train_path = os.path.join(settings.FILESPACE_ROOT_PATH, ai_train.train_code)
nested_train_path = os.path.join(
settings.FILESPACE_ROOT_PATH,
ai_algorithm_class.classcode,
ai_algorithm.algorithm_code,
"trains",
ai_train.train_code,
)
if getattr(ai_algorithm, "algorithm_engine", None) == "02":
candidates = [
os.path.join(nested_train_path, "output", f"{run_name}.log"),
os.path.join(flat_train_path, "output", f"{run_name}.log"),
]
elif getattr(ai_algorithm, "algorithm_engine", None) in {"03", "04"}:
candidates = [
os.path.join(nested_train_path, "runs", f"{run_name}.log"),
os.path.join(flat_train_path, "runs", f"{run_name}.log"),
]
else:
candidates = [
os.path.join(flat_train_path, "runs", f"{run_name}.log"),
os.path.join(nested_train_path, "runs", f"{run_name}.log"),
]
for path in candidates:
if os.path.exists(path):
return path
return candidates[0]
async def _emit_log_lines(self, payload: bytes) -> None:
for raw_line in payload.splitlines(keepends=False):
line = raw_line.decode("utf-8", errors="ignore").strip()
if not line:
continue
if self._is_valid_json(line):
await self.send(text_data=line)
else:
await self._send_json({"type": "waiting", "text": line})
async def _tail_train_log(self) -> None:
try:
log_file_path = await self._resolve_log_file_path()
except Exception as e:
await self._send_json({"type": "error", "text": f"无法定位训练日志文件:{str(e)}"})
return
last_pos = 0
buf = b""
while not self._stop_event.is_set():
if not os.path.exists(log_file_path):
last_pos = 0
buf = b""
await asyncio.sleep(0.5)
continue
try:
current_size = os.path.getsize(log_file_path)
if current_size < last_pos:
last_pos = 0
buf = b""
if current_size == last_pos:
await asyncio.sleep(0.2)
continue
with open(log_file_path, "rb") as f:
f.seek(last_pos)
chunk = f.read()
last_pos = f.tell()
if not chunk:
await asyncio.sleep(0.2)
continue
buf += chunk
normalized = buf.replace(b"\r\n", b"\n").replace(b"\r", b"\n")
if normalized.endswith(b"\n"):
completed = normalized
buf = b""
else:
split_idx = normalized.rfind(b"\n")
if split_idx == -1:
buf = normalized
await asyncio.sleep(0.2)
continue
completed = normalized[: split_idx + 1]
buf = normalized[split_idx + 1 :]
await self._emit_log_lines(completed)
await asyncio.sleep(0.05)
except asyncio.CancelledError:
return
except Exception as e:
await self._send_json({"type": "error", "text": f"训练日志读取失败:{str(e)}"})
await asyncio.sleep(1.0)