Skip to main content

weft_core/api/
package_ws.rs

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
15/// GET /ws/plugins/{package_name} -> WebSocket upgrade
16pub 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
48/// Reusable WASM payload dispatcher for plugins.
49/// This wrapper allows other modules to reuse the existing dispatch logic
50/// without duplicating behavior.
51pub 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}