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)