main
py 165 lines 6.17 KB
Raw
1 import threading
2
3 from flask import Flask
4
5 from helpers import cache, files, login, plugins, subagents
6 from helpers.api import CACHE_AREA, register_api_route
7 from helpers.extension import get_webui_extension_manifest
8 from helpers.ui_server import UiServerRuntime
9
10
11 WEBUI_MANIFEST_CACHE_AREA = "webui_extension_manifest(extensions)(plugins)"
12
13
14 def _new_app(name: str) -> Flask:
15 app = Flask(name, static_folder=None)
16 app.secret_key = "test-secret"
17 return app
18
19
20 def _api_handler_source(source: str) -> str:
21 return f"""from helpers.api import ApiHandler
22
23
24 class Handler(ApiHandler):
25 @classmethod
26 def get_methods(cls):
27 return ["GET"]
28
29 async def process(self, input, request):
30 return {{"source": {source!r}}}
31 """
32
33
34 def test_http_dispatches_contained_user_api_handler(tmp_path, monkeypatch) -> None:
35 monkeypatch.setattr(files, "_base_dir", str(tmp_path))
36 user_api_dir = tmp_path / "usr" / "api"
37 user_api_dir.mkdir(parents=True)
38 handler_source = _api_handler_source("user")
39 (user_api_dir / "ping.py").write_text(handler_source, encoding="utf-8")
40 (tmp_path / "usr" / "escaped.py").write_text(
41 handler_source, encoding="utf-8"
42 )
43 monkeypatch.setattr(login, "get_credentials_hash", lambda: "credential-hash")
44
45 cache.clear(CACHE_AREA)
46 try:
47 app = _new_app("test_user_api_route")
48 app.add_url_rule("/", "serve_index", lambda: "")
49 app.add_url_rule("/login", "login_handler", lambda: "")
50 register_api_route(app, threading.RLock())
51 client = app.test_client()
52
53 assert client.get("/api/ping").status_code == 302
54 with client.session_transaction() as session:
55 session["authentication"] = "credential-hash"
56 session["csrf_token"] = "csrf-token"
57 response = client.get("/api/ping", headers={"X-CSRF-Token": "csrf-token"})
58 assert response.status_code == 200
59 assert response.get_json() == {"source": "user"}
60
61 with app.test_request_context("/api/../escaped", method="GET"):
62 denied = app.ensure_sync(app.view_functions["api_dispatch"])("../escaped")
63 assert denied.status_code == 404
64 finally:
65 cache.clear(CACHE_AREA)
66
67
68 def test_existing_api_sources_keep_precedence(tmp_path, monkeypatch) -> None:
69 monkeypatch.setattr(files, "_base_dir", str(tmp_path))
70 monkeypatch.setattr(login, "get_credentials_hash", lambda: "credential-hash")
71
72 builtin_file = tmp_path / "api" / "shared.py"
73 builtin_file.parent.mkdir(parents=True)
74 builtin_file.write_text(_api_handler_source("builtin"), encoding="utf-8")
75
76 user_api_dir = tmp_path / "usr" / "api"
77 (user_api_dir / "plugins" / "demo").mkdir(parents=True)
78 (user_api_dir / "shared.py").write_text(
79 _api_handler_source("user"), encoding="utf-8"
80 )
81 (user_api_dir / "plugins" / "demo" / "ping.py").write_text(
82 _api_handler_source("user"), encoding="utf-8"
83 )
84
85 plugin_dir = tmp_path / "plugins" / "demo"
86 (plugin_dir / "api").mkdir(parents=True)
87 (plugin_dir / "api" / "ping.py").write_text(
88 _api_handler_source("plugin"), encoding="utf-8"
89 )
90 monkeypatch.setattr(
91 plugins,
92 "find_plugin_dir",
93 lambda name: str(plugin_dir) if name == "demo" else None,
94 )
95
96 cache.clear(CACHE_AREA)
97 try:
98 app = _new_app("test_existing_api_precedence")
99 app.add_url_rule("/", "serve_index", lambda: "")
100 app.add_url_rule("/login", "login_handler", lambda: "")
101 register_api_route(app, threading.RLock())
102 client = app.test_client()
103 with client.session_transaction() as session:
104 session["authentication"] = "credential-hash"
105 session["csrf_token"] = "csrf-token"
106 headers = {"X-CSRF-Token": "csrf-token"}
107
108 assert client.get("/api/shared", headers=headers).get_json() == {
109 "source": "builtin"
110 }
111 assert client.get("/api/plugins/demo/ping", headers=headers).get_json() == {
112 "source": "plugin"
113 }
114 finally:
115 cache.clear(CACHE_AREA)
116
117
118 def test_user_webui_manifest_asset_is_served_from_its_declared_url(
119 tmp_path, monkeypatch
120 ) -> None:
121 monkeypatch.setattr(files, "_base_dir", str(tmp_path))
122 extension_root = tmp_path / "usr" / "extensions" / "webui"
123 extension_file = extension_root / "route-probe" / "probe.js"
124 extension_file.parent.mkdir(parents=True)
125 extension_file.write_text("export default true;", encoding="utf-8")
126 builtin_extension_file = (
127 tmp_path / "extensions" / "webui" / "route-probe" / "probe.js"
128 )
129 builtin_extension_file.parent.mkdir(parents=True)
130 builtin_extension_file.write_text("export default false;", encoding="utf-8")
131 (extension_root.parent / "escaped.js").write_text("secret", encoding="utf-8")
132 monkeypatch.setattr(subagents, "get_paths", lambda *_args, **_kwargs: [str(extension_root)])
133
134 cache.clear(WEBUI_MANIFEST_CACHE_AREA)
135 try:
136 manifest = get_webui_extension_manifest(agent=None)
137 asset_url = manifest["js"]["route-probe"][0]
138 assert asset_url == "/usr/extensions/webui/route-probe/probe.js"
139
140 app = _new_app("test_user_webui_extension_route")
141 runtime = UiServerRuntime(
142 app, None, None, threading.RLock(), {} # type: ignore[arg-type]
143 )
144 runtime.register_http_routes()
145
146 client = app.test_client()
147 monkeypatch.setattr(login, "get_credentials_hash", lambda: "credential-hash")
148 assert client.get(asset_url).status_code == 302
149
150 monkeypatch.setattr(login, "get_credentials_hash", lambda: None)
151 builtin_response = client.get("/extensions/webui/route-probe/probe.js")
152 assert builtin_response.status_code == 200
153 assert builtin_response.get_data(as_text=True) == "export default false;"
154
155 response = client.get(asset_url)
156 assert response.status_code == 200
157 assert response.get_data(as_text=True) == "export default true;"
158
159 with app.test_request_context("/usr/extensions/webui/../escaped.js"):
160 denied = app.ensure_sync(
161 app.view_functions["serve_user_extension_asset"]
162 )("../escaped.js")
163 assert denied.status_code == 403
164 finally:
165 cache.clear(WEBUI_MANIFEST_CACHE_AREA)