From 4f976f6080489916312274873021b1f9f5671064 Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sun, 27 Oct 2024 17:11:20 +0000 Subject: [PATCH] Add /logs/raw and /logs/subscribe for getting logs on frontend Hijacks stderr/stdout to send all output data to the client on flush --- api_server/routes/internal/internal_routes.py | 29 ++++++++- api_server/services/terminal_service.py | 63 +++++++++++++++++++ app/logger.py | 62 ++++++++++++++---- server.py | 2 +- 4 files changed, 142 insertions(+), 14 deletions(-) create mode 100644 api_server/services/terminal_service.py diff --git a/api_server/routes/internal/internal_routes.py b/api_server/routes/internal/internal_routes.py index 63704f13a..3e0ad9193 100644 --- a/api_server/routes/internal/internal_routes.py +++ b/api_server/routes/internal/internal_routes.py @@ -2,6 +2,7 @@ from aiohttp import web from typing import Optional from folder_paths import models_dir, user_directory, output_directory, folder_names_and_paths from api_server.services.file_service import FileService +from api_server.services.terminal_service import TerminalService import app.logger class InternalRoutes: @@ -11,7 +12,8 @@ class InternalRoutes: Check README.md for more information. ''' - def __init__(self): + + def __init__(self, prompt_server): self.routes: web.RouteTableDef = web.RouteTableDef() self._app: Optional[web.Application] = None self.file_service = FileService({ @@ -19,6 +21,8 @@ class InternalRoutes: "user": user_directory, "output": output_directory }) + self.prompt_server = prompt_server + self.terminal_service = TerminalService(prompt_server) def setup_routes(self): @self.routes.get('/files') @@ -34,7 +38,28 @@ class InternalRoutes: @self.routes.get('/logs') async def get_logs(request): - return web.json_response(app.logger.get_logs()) + return web.json_response([(l["t"] + " - " + l["m"]) for l in app.logger.get_logs()]) + + @self.routes.get('/logs/raw') + async def get_logs(request): + self.terminal_service.update_size() + return web.json_response({ + "entries": list(app.logger.get_logs()), + "size": {"cols": self.terminal_service.cols, "rows": self.terminal_service.rows} + }) + + @self.routes.patch('/logs/subscribe') + async def subscribe_logs(request): + json_data = await request.json() + client_id = json_data["clientId"] + enabled = json_data["enabled"] + if enabled: + self.terminal_service.subscribe(client_id) + else: + self.terminal_service.unsubscribe(client_id) + + return web.Response(status=200) + @self.routes.get('/folder_paths') async def get_folder_paths(request): diff --git a/api_server/services/terminal_service.py b/api_server/services/terminal_service.py new file mode 100644 index 000000000..9701d1656 --- /dev/null +++ b/api_server/services/terminal_service.py @@ -0,0 +1,63 @@ +import asyncio +from app.logger import on_flush +import os + + +class TerminalService: + def __init__(self, server): + self.server = server + self.cols = None + self.rows = None + self.subscriptions = set() + on_flush(self.send_messages_sync) + + def update_size(self): + sz = os.get_terminal_size() + changed = False + if sz.columns != self.cols: + self.cols = sz.columns + changed = True + + if sz.lines != self.rows: + self.rows = sz.lines + changed = True + + if changed: + return {"cols": self.cols, "rows": self.rows} + + return None + + def subscribe(self, client_id): + self.subscriptions.add(client_id) + + def unsubscribe(self, client_id): + self.subscriptions.discard(client_id) + + def send_messages_sync(self, entries): + if not len(entries) or not len(self.subscriptions): + return + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + + if loop and loop.is_running(): + loop.create_task(self.send_messages(entries)) + else: + asyncio.run(self.send_messages(entries)) + + async def send_messages(self, entries): + if not len(entries): + return + + new_size = self.update_size() + + for client_id in self.subscriptions: + if client_id not in self.server.sockets: + # Automatically unsub if the socket has disconnected + self.unsubscribe(client_id) + continue + + await self.server.send_json( + "logs", {"entries": entries, "size": new_size}, client_id) diff --git a/app/logger.py b/app/logger.py index 4ca0ea88e..022e67297 100644 --- a/app/logger.py +++ b/app/logger.py @@ -1,20 +1,67 @@ -import logging -from logging.handlers import MemoryHandler from collections import deque +from datetime import datetime +import io +import logging +import sys +import threading logs = None -formatter = logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s") +stdout_interceptor = None +stderr_interceptor = None + + +class LogInterceptor(io.TextIOWrapper): + def __init__(self, stream, *args, **kwargs): + buffer = stream.buffer + encoding = stream.encoding + super().__init__(buffer, *args, **kwargs, encoding=encoding) + self._lock = threading.Lock() + self._flush_callbacks = [] + self._logs_since_flush = [] + + def write(self, data): + entry = {"t": datetime.now().isoformat(), "m": data} + with self._lock: + self._logs_since_flush.append(entry) + + # Simple handling for cr to overwrite the last output if it isnt a full line + # else logs just get full of progress messages + if isinstance(data, str) and data.startswith("\r") and not logs[-1]["m"].endswith("\n"): + logs.pop() + logs.append(entry) + super().write(data) + + def flush(self): + super().flush() + for cb in self._flush_callbacks: + cb(self._logs_since_flush) + self._logs_since_flush = [] + + def on_flush(self, callback): + self._flush_callbacks.append(callback) def get_logs(): - return "\n".join([formatter.format(x) for x in logs]) + return logs +def on_flush(callback): + stdout_interceptor.on_flush(callback) + stderr_interceptor.on_flush(callback) + def setup_logger(log_level: str = 'INFO', capacity: int = 300): global logs if logs: return + # Override output streams and log to buffer + logs = deque(maxlen=capacity) + + global stdout_interceptor + global stderr_interceptor + stdout_interceptor = sys.stdout = LogInterceptor(sys.stdout) + stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr) + # Setup default global logger logger = logging.getLogger() logger.setLevel(log_level) @@ -22,10 +69,3 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300): stream_handler = logging.StreamHandler() stream_handler.setFormatter(logging.Formatter("%(message)s")) logger.addHandler(stream_handler) - - # Create a memory handler with a deque as its buffer - logs = deque(maxlen=capacity) - memory_handler = MemoryHandler(capacity, flushLevel=logging.INFO) - memory_handler.buffer = logs - memory_handler.setFormatter(formatter) - logger.addHandler(memory_handler) diff --git a/server.py b/server.py index ada6d90c3..e663095bc 100644 --- a/server.py +++ b/server.py @@ -152,7 +152,7 @@ class PromptServer(): mimetypes.types_map['.js'] = 'application/javascript; charset=utf-8' self.user_manager = UserManager() - self.internal_routes = InternalRoutes() + self.internal_routes = InternalRoutes(self) self.supports = ["custom_nodes_from_web"] self.prompt_queue = None self.loop = loop