main
py 452 lines 15.7 KB
Raw
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)