| 1 | from dataclasses import dataclass, field |
| 2 | from datetime import timedelta |
| 3 | import asyncio |
| 4 | import gzip |
| 5 | import json |
| 6 | import logging |
| 7 | import os |
| 8 | import secrets |
| 9 | import threading |
| 10 | import time |
| 11 | from typing import Any |
| 12 | |
| 13 | from flask import ( |
| 14 | Flask, |
| 15 | Response, |
| 16 | redirect, |
| 17 | render_template_string, |
| 18 | request, |
| 19 | send_file, |
| 20 | session, |
| 21 | url_for, |
| 22 | ) |
| 23 | from socketio import ASGIApp |
| 24 | from starlette.applications import Starlette |
| 25 | from starlette.middleware.gzip import GZipMiddleware |
| 26 | from starlette.routing import Mount |
| 27 | from uvicorn.middleware.wsgi import WSGIMiddleware |
| 28 | from werkzeug.wrappers.request import Request as WerkzeugRequest |
| 29 | import socketio # type: ignore[import-untyped] |
| 30 | |
| 31 | from helpers import dotenv, fasta2a_server, files, git, login, mcp_server, runtime |
| 32 | from helpers.api import get_safe_next_url, register_api_route, requires_auth |
| 33 | from helpers.extension import extensible, get_webui_extension_manifest |
| 34 | from helpers.files import get_abs_path |
| 35 | from helpers.print_style import PrintStyle |
| 36 | from helpers.server_startup import StartupMonitor |
| 37 | from helpers.ui_bundler import ( |
| 38 | get_ui_asset_bundle, |
| 39 | serialize_ui_asset_bundle, |
| 40 | ) |
| 41 | from helpers import settings as settings_helper |
| 42 | from helpers.ws import register_ws_namespace, validate_ws_origin |
| 43 | from helpers.ws_manager import WsManager, set_shared_ws_manager |
| 44 | |
| 45 | |
| 46 | UPLOAD_LIMIT_BYTES = 5 * 1024 * 1024 * 1024 |
| 47 | SOCKETIO_PING_INTERVAL_SECONDS = 45 |
| 48 | SOCKETIO_PING_TIMEOUT_SECONDS = 120 |
| 49 | GZIP_MINIMUM_RESPONSE_BYTES = 1024 |
| 50 | GZIP_COMPRESSION_LEVEL = 6 |
| 51 | UI_INDEX_ASSET_URL = "/index.html" |
| 52 | |
| 53 | |
| 54 | def _positive_int_env(name: str, default: int) -> int: |
| 55 | raw_value = os.getenv(name) |
| 56 | if raw_value is None: |
| 57 | return default |
| 58 | try: |
| 59 | value = int(raw_value) |
| 60 | except (TypeError, ValueError): |
| 61 | return default |
| 62 | return value if value > 0 else default |
| 63 | |
| 64 | |
| 65 | def configure_process_environment() -> None: |
| 66 | logging.getLogger().setLevel(logging.WARNING) |
| 67 | os.environ["TOKENIZERS_PARALLELISM"] = "false" |
| 68 | from helpers.localization import Localization |
| 69 | |
| 70 | Localization.get().apply_process_timezone() |
| 71 | |
| 72 | |
| 73 | @dataclass |
| 74 | class UiServerRuntime: |
| 75 | webapp: Flask |
| 76 | socketio_server: socketio.AsyncServer |
| 77 | ws_manager: WsManager |
| 78 | lock: threading.RLock |
| 79 | settings_snapshot: dict[str, Any] |
| 80 | _routes_registered: bool = False |
| 81 | _transport_registered: bool = False |
| 82 | _route_handlers: "UiRouteHandlers | None" = field(default=None, init=False) |
| 83 | |
| 84 | @classmethod |
| 85 | def create(cls) -> "UiServerRuntime": |
| 86 | webapp = Flask("app", static_folder=get_abs_path("./webui"), static_url_path="/") |
| 87 | webapp.secret_key = os.getenv("FLASK_SECRET_KEY") or secrets.token_hex(32) |
| 88 | |
| 89 | WerkzeugRequest.max_form_memory_size = UPLOAD_LIMIT_BYTES |
| 90 | webapp.config.update( |
| 91 | JSON_SORT_KEYS=False, |
| 92 | SESSION_COOKIE_NAME="session_" + runtime.get_runtime_id(), |
| 93 | SESSION_COOKIE_SAMESITE="Lax", |
| 94 | SESSION_PERMANENT=True, |
| 95 | PERMANENT_SESSION_LIFETIME=timedelta(days=1), |
| 96 | MAX_CONTENT_LENGTH=int( |
| 97 | os.getenv("FLASK_MAX_CONTENT_LENGTH", str(UPLOAD_LIMIT_BYTES)) |
| 98 | ), |
| 99 | MAX_FORM_MEMORY_SIZE=int( |
| 100 | os.getenv("FLASK_MAX_FORM_MEMORY_SIZE", str(UPLOAD_LIMIT_BYTES)) |
| 101 | ), |
| 102 | ) |
| 103 | |
| 104 | lock = threading.RLock() |
| 105 | socketio_server = socketio.AsyncServer( |
| 106 | async_mode="asgi", |
| 107 | namespaces="*", |
| 108 | cors_allowed_origins=lambda _origin, environ: validate_ws_origin(environ)[0], |
| 109 | logger=False, |
| 110 | engineio_logger=False, |
| 111 | ping_interval=_positive_int_env( |
| 112 | "A0_SOCKETIO_PING_INTERVAL_SECONDS", |
| 113 | SOCKETIO_PING_INTERVAL_SECONDS, |
| 114 | ), |
| 115 | ping_timeout=_positive_int_env( |
| 116 | "A0_SOCKETIO_PING_TIMEOUT_SECONDS", |
| 117 | SOCKETIO_PING_TIMEOUT_SECONDS, |
| 118 | ), |
| 119 | max_http_buffer_size=50 * 1024 * 1024, |
| 120 | ) |
| 121 | |
| 122 | ws_manager = WsManager(socketio_server, lock) |
| 123 | set_shared_ws_manager(ws_manager) |
| 124 | |
| 125 | server_runtime = cls( |
| 126 | webapp=webapp, |
| 127 | socketio_server=socketio_server, |
| 128 | ws_manager=ws_manager, |
| 129 | lock=lock, |
| 130 | settings_snapshot={}, |
| 131 | ) |
| 132 | server_runtime.refresh_runtime_settings() |
| 133 | return server_runtime |
| 134 | |
| 135 | def refresh_runtime_settings(self) -> None: |
| 136 | self.settings_snapshot = settings_helper.get_settings() |
| 137 | settings_helper.set_runtime_settings_snapshot(self.settings_snapshot) |
| 138 | self.ws_manager.set_server_restart_broadcast( |
| 139 | self.settings_snapshot.get("websocket_server_restart_enabled", True) |
| 140 | ) |
| 141 | |
| 142 | def register_http_routes(self) -> None: |
| 143 | if self._routes_registered: |
| 144 | return |
| 145 | |
| 146 | handlers = UiRouteHandlers(self) |
| 147 | self._route_handlers = handlers |
| 148 | self.webapp.add_url_rule( |
| 149 | "/login", |
| 150 | "login_handler", |
| 151 | handlers.login_handler, |
| 152 | methods=["GET", "POST"], |
| 153 | ) |
| 154 | self.webapp.add_url_rule( |
| 155 | "/logout", |
| 156 | "logout_handler", |
| 157 | handlers.logout_handler, |
| 158 | methods=["GET"], |
| 159 | ) |
| 160 | self.webapp.add_url_rule( |
| 161 | "/", |
| 162 | "serve_index", |
| 163 | handlers.serve_splash, |
| 164 | methods=["GET"], |
| 165 | ) |
| 166 | self.webapp.add_url_rule( |
| 167 | "/index.html", |
| 168 | "serve_app_index", |
| 169 | handlers.serve_index, |
| 170 | methods=["GET"], |
| 171 | ) |
| 172 | self.webapp.add_url_rule( |
| 173 | "/ui/index", |
| 174 | "serve_bootstrap_index", |
| 175 | handlers.serve_index, |
| 176 | methods=["GET"], |
| 177 | ) |
| 178 | self.webapp.add_url_rule( |
| 179 | "/safe", |
| 180 | "serve_safe", |
| 181 | handlers.serve_safe, |
| 182 | methods=["GET"], |
| 183 | ) |
| 184 | self.webapp.add_url_rule( |
| 185 | "/ui/asset-bundle", |
| 186 | "serve_ui_asset_bundle", |
| 187 | handlers.serve_ui_asset_bundle, |
| 188 | methods=["GET"], |
| 189 | ) |
| 190 | self.webapp.add_url_rule( |
| 191 | "/plugins/<plugin_name>/<path:asset_path>", |
| 192 | "serve_builtin_plugin_asset", |
| 193 | handlers.serve_builtin_plugin_asset, |
| 194 | methods=["GET"], |
| 195 | ) |
| 196 | self.webapp.add_url_rule( |
| 197 | "/usr/plugins/<plugin_name>/<path:asset_path>", |
| 198 | "serve_plugin_asset", |
| 199 | handlers.serve_plugin_asset, |
| 200 | methods=["GET"], |
| 201 | ) |
| 202 | self.webapp.add_url_rule( |
| 203 | "/extensions/webui/<path:asset_path>", |
| 204 | "serve_extension_asset", |
| 205 | handlers.serve_extension_asset, |
| 206 | methods=["GET"], |
| 207 | ) |
| 208 | self.webapp.add_url_rule( |
| 209 | "/usr/extensions/webui/<path:asset_path>", |
| 210 | "serve_user_extension_asset", |
| 211 | handlers.serve_user_extension_asset, |
| 212 | methods=["GET"], |
| 213 | ) |
| 214 | self._routes_registered = True |
| 215 | |
| 216 | def register_transport_handlers(self) -> None: |
| 217 | if self._transport_registered: |
| 218 | return |
| 219 | register_api_route(self.webapp, self.lock) |
| 220 | register_ws_namespace( |
| 221 | self.socketio_server, |
| 222 | self.webapp, |
| 223 | self.lock, |
| 224 | manager=self.ws_manager, |
| 225 | ) |
| 226 | self._transport_registered = True |
| 227 | |
| 228 | def build_asgi_app(self, startup_monitor: StartupMonitor): |
| 229 | with startup_monitor.stage("wsgi.middleware.create"): |
| 230 | wsgi_app = WSGIMiddleware(self.webapp) |
| 231 | |
| 232 | with startup_monitor.stage("mcp.proxy.init"): |
| 233 | mcp_app = mcp_server.DynamicMcpProxy.get_instance() |
| 234 | |
| 235 | with startup_monitor.stage("a2a.proxy.init"): |
| 236 | a2a_app = fasta2a_server.DynamicA2AProxy.get_instance() |
| 237 | |
| 238 | with startup_monitor.stage("starlette.app.create"): |
| 239 | starlette_app = Starlette( |
| 240 | routes=[ |
| 241 | Mount("/mcp", app=mcp_app), |
| 242 | Mount("/a2a", app=a2a_app), |
| 243 | Mount("/", app=wsgi_app), |
| 244 | ], |
| 245 | lifespan=startup_monitor.lifespan(), |
| 246 | ) |
| 247 | compressed_http_app = GZipMiddleware( |
| 248 | starlette_app, |
| 249 | minimum_size=GZIP_MINIMUM_RESPONSE_BYTES, |
| 250 | compresslevel=GZIP_COMPRESSION_LEVEL, |
| 251 | ) |
| 252 | |
| 253 | with startup_monitor.stage("socketio.asgi.create"): |
| 254 | return ASGIApp(self.socketio_server, other_asgi_app=compressed_http_app) |
| 255 | |
| 256 | def access_log_enabled(self) -> bool: |
| 257 | return self.settings_snapshot.get("uvicorn_access_logs_enabled", False) |
| 258 | |
| 259 | |
| 260 | class UiRouteHandlers: |
| 261 | def __init__(self, runtime_state: UiServerRuntime) -> None: |
| 262 | self.runtime = runtime_state |
| 263 | |
| 264 | @extensible |
| 265 | async def login_handler(self): |
| 266 | error = None |
| 267 | fallback_url = url_for("serve_index") |
| 268 | next_url = get_safe_next_url( |
| 269 | request.form.get("next") if request.method == "POST" else request.args.get("next"), |
| 270 | fallback_url, |
| 271 | ) |
| 272 | |
| 273 | if request.method == "POST": |
| 274 | user = dotenv.get_dotenv_value("AUTH_LOGIN") |
| 275 | password = dotenv.get_dotenv_value("AUTH_PASSWORD") |
| 276 | |
| 277 | if request.form["username"] == user and request.form["password"] == password: |
| 278 | session["authentication"] = login.get_credentials_hash() |
| 279 | return redirect(next_url or fallback_url) |
| 280 | else: |
| 281 | await asyncio.sleep(1) |
| 282 | error = "Invalid Credentials. Please try again." |
| 283 | |
| 284 | login_page_content = files.read_file("webui/login.html") |
| 285 | return render_template_string(login_page_content, error=error, next=next_url) |
| 286 | |
| 287 | @extensible |
| 288 | async def logout_handler(self): |
| 289 | session.pop("authentication", None) |
| 290 | return redirect(url_for("login_handler")) |
| 291 | |
| 292 | @requires_auth |
| 293 | async def serve_splash(self): |
| 294 | return Response( |
| 295 | files.read_file("webui/splash.html"), |
| 296 | content_type="text/html; charset=utf-8", |
| 297 | headers={"Cache-Control": "no-store"}, |
| 298 | ) |
| 299 | |
| 300 | @requires_auth |
| 301 | async def serve_safe(self): |
| 302 | if request.args.get("__direct") == "1": |
| 303 | return await self.serve_index() |
| 304 | return Response( |
| 305 | files.read_file("webui/safe.html"), |
| 306 | content_type="text/html; charset=utf-8", |
| 307 | headers={"Cache-Control": "no-store"}, |
| 308 | ) |
| 309 | |
| 310 | @requires_auth |
| 311 | @extensible |
| 312 | async def serve_index(self): |
| 313 | try: |
| 314 | gitinfo = git.get_git_info() |
| 315 | except Exception: |
| 316 | gitinfo = { |
| 317 | "version": "unknown", |
| 318 | "commit_time": "unknown", |
| 319 | } |
| 320 | try: |
| 321 | user_timezone_setting = str(settings_helper.get_settings().get("timezone", "auto")) |
| 322 | except Exception: |
| 323 | user_timezone_setting = "auto" |
| 324 | try: |
| 325 | user_time_format_setting = str(settings_helper.get_settings().get("time_format", "12h")) |
| 326 | except Exception: |
| 327 | user_time_format_setting = "12h" |
| 328 | try: |
| 329 | user_ui_control_visibility = json.dumps( |
| 330 | settings_helper.get_settings()["ui_control_visibility"], |
| 331 | separators=(",", ":"), |
| 332 | ) |
| 333 | except Exception: |
| 334 | user_ui_control_visibility = json.dumps(settings_helper.UI_CONTROL_VISIBILITY_DEFAULTS) |
| 335 | try: |
| 336 | webui_extension_manifest = json.dumps( |
| 337 | get_webui_extension_manifest(agent=None), |
| 338 | separators=(",", ":"), |
| 339 | ) |
| 340 | webui_extension_manifest = ( |
| 341 | webui_extension_manifest.replace("&", "\\u0026") |
| 342 | .replace("<", "\\u003c") |
| 343 | .replace(">", "\\u003e") |
| 344 | ) |
| 345 | except Exception: |
| 346 | webui_extension_manifest = "null" |
| 347 | |
| 348 | index = files.read_file("webui/index.html") |
| 349 | return files.replace_placeholders_text( |
| 350 | _content=index, |
| 351 | version_no=gitinfo["version"], |
| 352 | version_time=gitinfo["commit_time"], |
| 353 | runtime_id=runtime.get_runtime_id(), |
| 354 | runtime_is_development=("true" if runtime.is_development() else "false"), |
| 355 | logged_in=("true" if login.get_credentials_hash() else "false"), |
| 356 | user_timezone_setting=user_timezone_setting, |
| 357 | user_time_format_setting=user_time_format_setting, |
| 358 | user_ui_control_visibility=user_ui_control_visibility, |
| 359 | webui_extension_manifest=webui_extension_manifest, |
| 360 | ) |
| 361 | |
| 362 | @requires_auth |
| 363 | async def serve_ui_asset_bundle(self): |
| 364 | try: |
| 365 | bundle = get_ui_asset_bundle([UI_INDEX_ASSET_URL], agent=None) |
| 366 | return self._serve_ui_asset_payload(bundle) |
| 367 | except Exception as error: |
| 368 | PrintStyle.warning(f"Unable to build WebUI asset bundle: {error}") |
| 369 | return Response( |
| 370 | '{"error":"WebUI asset bundle unavailable"}', |
| 371 | status=503, |
| 372 | content_type="application/json; charset=utf-8", |
| 373 | headers={"Cache-Control": "no-store"}, |
| 374 | ) |
| 375 | |
| 376 | def _serve_ui_asset_payload(self, asset_payload: dict): |
| 377 | version = str(asset_payload.get("version") or "") |
| 378 | if not version: |
| 379 | raise ValueError("WebUI asset payload has no version") |
| 380 | if request.if_none_match.contains_weak(version): |
| 381 | response = Response(status=304) |
| 382 | response.headers["Vary"] = "Accept-Encoding" |
| 383 | response.set_etag(version, weak=True) |
| 384 | response.cache_control.private = True |
| 385 | response.cache_control.no_cache = True |
| 386 | return response |
| 387 | |
| 388 | payload = serialize_ui_asset_bundle(asset_payload).encode("utf-8") |
| 389 | use_gzip = request.accept_encodings["gzip"] > 0 |
| 390 | response = Response( |
| 391 | gzip.compress(payload) if use_gzip else payload, |
| 392 | content_type="application/json; charset=utf-8", |
| 393 | ) |
| 394 | if use_gzip: |
| 395 | response.headers["Content-Encoding"] = "gzip" |
| 396 | response.headers["Vary"] = "Accept-Encoding" |
| 397 | response.set_etag(version, weak=True) |
| 398 | response.cache_control.private = True |
| 399 | response.cache_control.no_cache = True |
| 400 | return response |
| 401 | |
| 402 | @requires_auth |
| 403 | async def serve_builtin_plugin_asset(self, plugin_name, asset_path): |
| 404 | return await self._serve_plugin_asset(plugin_name, asset_path) |
| 405 | |
| 406 | @requires_auth |
| 407 | async def serve_plugin_asset(self, plugin_name, asset_path): |
| 408 | return await self._serve_plugin_asset(plugin_name, asset_path) |
| 409 | |
| 410 | @requires_auth |
| 411 | async def serve_extension_asset(self, asset_path): |
| 412 | return self._serve_extension_asset( |
| 413 | files.get_abs_path("extensions/webui"), asset_path |
| 414 | ) |
| 415 | |
| 416 | @requires_auth |
| 417 | async def serve_user_extension_asset(self, asset_path): |
| 418 | return self._serve_extension_asset( |
| 419 | files.get_abs_path(files.USER_DIR, "extensions/webui"), asset_path |
| 420 | ) |
| 421 | |
| 422 | def _serve_extension_asset(self, extension_dir, asset_path): |
| 423 | path = files.get_abs_path(extension_dir, asset_path) |
| 424 | if not files.is_in_dir(path, extension_dir): |
| 425 | return Response("Access denied", 403) |
| 426 | return send_file(path) |
| 427 | |
| 428 | @extensible |
| 429 | async def _serve_plugin_asset(self, plugin_name, asset_path): |
| 430 | from helpers import plugins |
| 431 | |
| 432 | plugin_dir = plugins.find_plugin_dir(plugin_name) |
| 433 | if not plugin_dir: |
| 434 | return Response("Plugin not found", 404) |
| 435 | |
| 436 | try: |
| 437 | asset_file = files.get_abs_path(plugin_dir, asset_path) |
| 438 | webui_dir = files.get_abs_path(plugin_dir, "webui") |
| 439 | webui_extensions_dir = files.get_abs_path(plugin_dir, "extensions/webui") |
| 440 | |
| 441 | if not files.is_in_dir(str(asset_file), str(webui_dir)) and not files.is_in_dir( |
| 442 | str(asset_file), str(webui_extensions_dir) |
| 443 | ): |
| 444 | return Response("Access denied", 403) |
| 445 | |
| 446 | if not files.is_file(asset_file): |
| 447 | return Response("Asset not found", 404) |
| 448 | |
| 449 | return send_file(str(asset_file)) |
| 450 | except Exception as e: |
| 451 | PrintStyle.error(f"Error serving plugin asset: {e}") |
| 452 | return Response("Error serving asset", 500) |