1use crate::api::openai_compat::AppState;
2use crate::package::bridge::{resolve_loaded_package_name_from_aliases, WasmHandle};
3use crate::package::resolve_runtime_package;
4use axum::{
5 extract::{
6 ws::{Message, WebSocket, WebSocketUpgrade},
7 Path, Query, State,
8 },
9 http::StatusCode,
10 response::{IntoResponse, Response},
11 Json,
12};
13use std::collections::HashMap;
14
15pub async fn package_websocket(
17 ws: WebSocketUpgrade,
18 Path(package_name): Path<String>,
19 Query(params): Query<HashMap<String, String>>,
20 State(state): State<AppState>,
21) -> Response {
22 let token = state.runtime_token.as_deref().unwrap_or_default();
23 let authorized = params
24 .get("token")
25 .map(|value| value == token)
26 .unwrap_or(false);
27 if !authorized {
28 return (
29 StatusCode::UNAUTHORIZED,
30 Json(serde_json::json!({
31 "error": "missing or invalid loopback bearer token"
32 })),
33 )
34 .into_response();
35 }
36
37 ws.on_upgrade(move |socket| handle_plugin_ws(socket, package_name, state))
38}
39
40pub async fn package_call(
41 Path(package_name): Path<String>,
42 State(state): State<AppState>,
43 Json(payload): Json<serde_json::Value>,
44) -> impl IntoResponse {
45 Json(dispatch_package_payload(&package_name, payload, &state).await)
46}
47
48pub async fn dispatch_to_wasm_package_payload(
52 package_name: &str,
53 payload: serde_json::Value,
54 state: &AppState,
55) -> serde_json::Value {
56 let wasm_handle = state.wasm_handle.read().await.clone();
57 let canonical_package_name =
58 resolve_runtime_package(&state.repo_root, &state.package_index, package_name)
59 .map(|package| package.manifest.package_info.name)
60 .filter(|resolved| !resolved.trim().is_empty());
61
62 let target_package_name = if let Some(handle) = wasm_handle.as_ref() {
63 let package_aliases = {
64 let config = state.config.read().await;
65 config.package_aliases.clone()
66 };
67 let mut merged_aliases =
68 crate::package::merged_package_aliases(&state.package_index, &package_aliases);
69 if let Some(canonical) = canonical_package_name.as_ref() {
70 if canonical != package_name {
71 merged_aliases.insert(package_name.trim().to_string(), canonical.clone());
72 }
73 }
74 let loaded_package_names = handle.package_names();
75 Some(resolve_loaded_package_name_from_aliases(
76 &merged_aliases,
77 &loaded_package_names,
78 canonical_package_name.as_deref().unwrap_or(package_name),
79 ))
80 } else {
81 canonical_package_name.clone()
82 };
83
84 dispatch_to_plugin(
85 package_name,
86 target_package_name.as_deref(),
87 &payload,
88 &wasm_handle,
89 )
90 .await
91}
92
93async fn handle_plugin_ws(mut socket: WebSocket, package_name: String, state: AppState) {
94 tracing::info!("[ws] Package '{}' UI connected", package_name);
95
96 while let Some(Ok(msg)) = socket.recv().await {
97 match msg {
98 Message::Text(text) => {
99 let ws_msg: serde_json::Value = match serde_json::from_str(&text) {
100 Ok(m) => m,
101 Err(e) => {
102 let err = serde_json::json!({
103 "id": "error",
104 "type": "error",
105 "payload": {"message": format!("Invalid message: {}", e)}
106 });
107 let _ = socket.send(Message::Text(err.to_string().into())).await;
108 continue;
109 }
110 };
111
112 let message_id = ws_msg
113 .get("id")
114 .cloned()
115 .unwrap_or_else(|| serde_json::json!(""));
116 let response_payload = dispatch_package_payload(
117 &package_name,
118 ws_msg
119 .get("payload")
120 .cloned()
121 .unwrap_or_else(|| serde_json::json!({})),
122 &state,
123 )
124 .await;
125
126 let response = serde_json::json!({
127 "id": message_id,
128 "type": "response",
129 "payload": response_payload
130 });
131 let _ = socket
132 .send(Message::Text(response.to_string().into()))
133 .await;
134 }
135 Message::Close(_) => {
136 tracing::info!("[ws] Package '{}' UI disconnected", package_name);
137 break;
138 }
139 _ => {}
140 }
141 }
142}
143
144pub(crate) async fn dispatch_package_payload(
145 package_name: &str,
146 payload: serde_json::Value,
147 state: &AppState,
148) -> serde_json::Value {
149 if let Some(service_payload) = dispatch_to_service_package(package_name, &payload, state).await
150 {
151 return service_payload;
152 }
153
154 dispatch_to_wasm_package_payload(package_name, payload.clone(), state).await
155}
156
157async fn dispatch_to_service_package(
158 package_name: &str,
159 payload: &serde_json::Value,
160 state: &AppState,
161) -> Option<serde_json::Value> {
162 let resolved_package_name =
163 resolve_runtime_package(&state.repo_root, &state.package_index, package_name)
164 .map(|package| package.manifest.package_info.name)
165 .unwrap_or_else(|| package_name.to_string());
166
167 let service_config = state
168 .process_manager
169 .service_config(&resolved_package_name)
170 .await?;
171 let Some(health_url) = service_config.health_url.clone() else {
172 return Some(
173 serde_json::json!({"error": format!("Service package '{}' has no health url", resolved_package_name)}),
174 );
175 };
176
177 let base_url = health_url.trim_end_matches("/health");
178 let request_url = service_plugin_request_url(&resolved_package_name, base_url);
179 let client = reqwest::Client::new();
180
181 let request_body = service_plugin_request_body(&resolved_package_name, payload);
182
183 match client.post(&request_url).json(&request_body).send().await {
184 Ok(response) => {
185 let status = response.status();
186 match response.json::<serde_json::Value>().await {
187 Ok(body) => {
188 if status.is_success() {
189 Some(body)
190 } else {
191 Some(serde_json::json!({
192 "error": format!("service package '{}' returned HTTP {}", resolved_package_name, status),
193 "details": body,
194 }))
195 }
196 }
197 Err(error) => Some(serde_json::json!({
198 "error": format!("service package '{}' returned invalid JSON: {}", resolved_package_name, error),
199 })),
200 }
201 }
202 Err(error) => Some(serde_json::json!({
203 "error": format!("service package '{}' request failed: {}", resolved_package_name, error),
204 })),
205 }
206}
207
208fn service_plugin_request_url(package_name: &str, base_url: &str) -> String {
209 match package_name {
210 "js-extension-runtime" => format!("{}/execute", base_url),
211 _ => format!("{}/webhook", base_url),
212 }
213}
214
215fn service_plugin_request_body(
216 package_name: &str,
217 payload: &serde_json::Value,
218) -> serde_json::Value {
219 if package_name == "js-extension-runtime" {
220 let action = payload
221 .get("action")
222 .and_then(|value| value.as_str())
223 .unwrap_or("");
224 let data = payload
225 .get("data")
226 .cloned()
227 .unwrap_or_else(|| serde_json::json!({}));
228 let tool = match action {
229 "web_search" | "search_web" | "websearch" => "web_search",
230 "fetch_url" | "web_fetch" => "fetch_url",
231 _ => action,
232 };
233 let js_payload = if data.get("tool").is_some() || data.get("args").is_some() {
234 data
235 } else {
236 serde_json::json!({
237 "tool": tool,
238 "args": data,
239 })
240 };
241 return serde_json::json!({
242 "id": "weft-tool-executor",
243 "action": "execute_tool",
244 "payload": js_payload,
245 });
246 }
247 payload.clone()
248}
249async fn dispatch_to_plugin(
250 requested_package_name: &str,
251 resolved_package_name: Option<&str>,
252 payload: &serde_json::Value,
253 wasm_handle: &Option<WasmHandle>,
254) -> serde_json::Value {
255 let Some(handle) = wasm_handle else {
256 return serde_json::json!({"error": "No WASM runtime available"});
257 };
258
259 let target_package_name = resolved_package_name
260 .filter(|name| handle.has_package(name))
261 .unwrap_or(requested_package_name);
262
263 if !handle.has_package(target_package_name) {
264 return serde_json::json!({"error": format!("Package '{}' not loaded", target_package_name)});
265 }
266
267 let payload_str = serde_json::to_string(payload).unwrap_or_default();
268
269 let package_name = target_package_name.to_string();
270 let package_name_for_call = package_name.clone();
271 let handle = handle.clone();
272 let call_result = tokio::task::spawn_blocking(move || {
273 handle.call(&package_name_for_call, "handle_ws_message", &payload_str)
274 })
275 .await;
276
277 match call_result {
278 Ok(Ok(result_str)) => serde_json::from_str(&result_str)
279 .unwrap_or_else(|_| serde_json::json!({"result": result_str})),
280 Ok(Err(e)) => {
281 tracing::error!(
282 "[ws] Package '{}' handle_ws_message error: {}",
283 package_name,
284 e
285 );
286 serde_json::json!({"error": format!("{}", e)})
287 }
288 Err(e) => {
289 tracing::error!(
290 "[ws] Package '{}' join error while handling package call: {}",
291 package_name,
292 e
293 );
294 serde_json::json!({"error": format!("package worker join error: {}", e)})
295 }
296 }
297}
298
299#[cfg(test)]
300mod tests {
301 use super::{service_plugin_request_body, service_plugin_request_url};
302
303 #[test]
304 fn service_package_mapping_keeps_memory_runtime_on_webhook_passthrough() {
305 assert_eq!(
306 service_plugin_request_url("memory-runtime", "http://127.0.0.1:17830"),
307 "http://127.0.0.1:17830/webhook"
308 );
309
310 let payload = serde_json::json!({
311 "action": "write",
312 "data": {
313 "key": "session",
314 "value": "hello"
315 }
316 });
317
318 assert_eq!(
319 service_plugin_request_body("memory-runtime", &payload),
320 payload
321 );
322 }
323
324 #[test]
325 fn generic_service_http_mapping_uses_webhook_and_passthrough_body() {
326 let payload = serde_json::json!({"hello": "world"});
327
328 assert_eq!(
329 service_plugin_request_url("companion-core", "http://127.0.0.1:17830"),
330 "http://127.0.0.1:17830/webhook"
331 );
332 assert_eq!(
333 service_plugin_request_body("companion-core", &payload),
334 payload
335 );
336 }
337}