model_train_dm/backend/apps/training/consumers.py

151 lines
5.5 KiB
Python
Raw Normal View History

2026-07-27 17:51:49 +08:00
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)