Skip to main content

zerolaunch_plugin_sdk_rust/
host_proxy.rs

1//! HostProxy — 第三方插件调用 host/* API 的代理。
2//!
3//! 每个方法经共享出站通道发送 LSP 帧 JSON-RPC 请求,并通过注册在共享
4//! pending map 中的 oneshot 等待响应。该设计将 stdin 读取与 stdout 写入
5//! 集中到独立任务,避免旧同步 stdin 锁方案的死锁。
6use dashmap::DashMap;
7use std::sync::atomic::{AtomicU64, Ordering};
8use std::sync::Arc;
9use std::time::Duration;
10use tokio::sync::{mpsc, oneshot};
11use zerolaunch_plugin_api::services::model::{
12    ModelChatRequest, ModelChatResponse, ModelEmbeddingRequest, ModelEmbeddingResponse, ModelInfo,
13    ModelSimilarityRequest, ModelSimilarityResponse,
14};
15use zerolaunch_plugin_protocol::methods::host;
16use zerolaunch_plugin_protocol::JsonRpcError;
17
18use base64::Engine as _;
19
20/// Proxy for calling host-side APIs from a plugin subprocess.
21/// Does NOT access stdin/stdout directly — uses channel-based I/O.
22pub struct HostProxy {
23    /// 请求 id 分配器(单调递增,从 1 开始)。
24    next_id: AtomicU64,
25    /// 在途请求表:id → 响应通道。read_task 完成响应后移除条目。
26    /// 值类型为 `Result`:宿主错误经 Err(JsonRpcError) 返回,调用方区分错误与结果。
27    pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>>,
28    /// 出站帧通道:write_task 独占消费并写 stdout。
29    outbound_tx: mpsc::Sender<Vec<u8>>,
30}
31
32impl HostProxy {
33    pub fn new(
34        pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>>,
35        outbound_tx: mpsc::Sender<Vec<u8>>,
36    ) -> Self {
37        Self {
38            next_id: AtomicU64::new(1),
39            pending,
40            outbound_tx,
41        }
42    }
43
44    /// Send a host/* request via the shared stdout channel and await the response
45    /// through the shared pending map.
46    async fn send_request(
47        &self,
48        method: &str,
49        params: serde_json::Value,
50    ) -> Result<serde_json::Value, String> {
51        // 模型类方法耗时可能远超普通 host/*(LLM 生成 / 大批量 embedding),
52        // 单独放宽超时;其余方法维持 30s。
53        let is_model_call = method.starts_with("host/model.");
54        let timeout = if is_model_call {
55            Duration::from_secs(300)
56        } else {
57            Duration::from_secs(30)
58        };
59        self.send_request_inner(method, params, timeout).await
60    }
61
62    /// 发送请求并等待响应;超时/通道关闭返回错误字符串。
63    async fn send_request_inner(
64        &self,
65        method: &str,
66        params: serde_json::Value,
67        timeout: Duration,
68    ) -> Result<serde_json::Value, String> {
69        let id = self.next_id.fetch_add(1, Ordering::SeqCst);
70        let request = serde_json::json!({
71            "jsonrpc": "2.0",
72            "id": id,
73            "method": method,
74            "params": params,
75        });
76
77        let payload = serde_json::to_vec(&request).map_err(|e| e.to_string())?;
78
79        // Register pending response
80        let (tx, rx) = oneshot::channel();
81        self.pending.insert(id, tx);
82
83        // Send via the shared channel (write_task writes to stdout).
84        // 注意:此处只投递干净 JSON,分帧统一由 write_task 完成,
85        // 若在此 encode_frame 会导致双重分帧(帧内嵌套帧),宿主无法解析。
86        if self.outbound_tx.send(payload).await.is_err() {
87            self.pending.remove(&id);
88            return Err("write channel closed".to_string());
89        }
90
91        // Await the response (read_task completes the oneshot with Result).
92        // Apply a timeout so the plugin doesn't hang forever if
93        // the host crashes during request processing.
94        match tokio::time::timeout(timeout, rx).await {
95            Ok(Ok(Ok(value))) => Ok(value),
96            Ok(Ok(Err(err))) => Err(format!(
97                "host call failed (code {}): {}",
98                err.code, err.message
99            )),
100            Ok(Err(_)) => Err("response channel closed".to_string()),
101            Err(_) => {
102                // 超时:移除 pending 条目,防止长驻插件无界泄漏。
103                self.pending.remove(&id);
104                Err("host call timed out".to_string())
105            }
106        }
107    }
108
109    pub async fn log(&self, level: &str, message: &str) -> Result<(), String> {
110        self.send_request(
111            host::LOG,
112            serde_json::json!({ "level": level, "message": message }),
113        )
114        .await?;
115        Ok(())
116    }
117
118    /// 发送 host/log 请求但不等待响应(fire-and-forget)。
119    ///
120    /// pending 条目在宿主响应到达时由 read_task 自动清理。
121    /// 若 outbound 通道已满,日志被静默丢弃并从 pending 中移除,
122    /// 避免阻塞调用者(通常来自 tracing subscriber 的回调)。
123    pub fn log_no_wait(&self, level: &str, message: &str) {
124        let id = self.next_id.fetch_add(1, Ordering::SeqCst);
125        let Ok(payload) = serde_json::to_vec(&serde_json::json!({
126            "jsonrpc": "2.0",
127            "id": id,
128            "method": host::LOG,
129            "params": { "level": level, "message": message },
130        })) else {
131            return;
132        };
133
134        let (tx, _rx) = oneshot::channel(); // _rx 立即 drop → fire-and-forget
135        self.pending.insert(id, tx);
136
137        // 非阻塞投递:通道满了则丢弃并清理 pending。
138        // 同 send_request:只投递干净 JSON,分帧由 write_task 统一完成。
139        if self.outbound_tx.try_send(payload).is_err() {
140            self.pending.remove(&id);
141        }
142    }
143
144    pub async fn shell_open(&self, target: &str) -> Result<(), String> {
145        self.send_request(host::SHELL_OPEN, serde_json::json!({ "target": target }))
146            .await?;
147        Ok(())
148    }
149
150    /// 获取图标字节(base64 字符串)。
151    /// 返回:图标字节的 base64(WebP,回退可能为 PNG);失败为空字符串。
152    pub async fn get_icon(&self, path: &str) -> Result<String, String> {
153        let result = self
154            .send_request(
155                host::ICON_GET,
156                serde_json::json!({ "request": { "path": path }, "level": "Full" }),
157            )
158            .await?;
159        Ok(result.as_str().unwrap_or("").to_string())
160    }
161
162    pub async fn shell_execute_command(&self, cmd: &str) -> Result<(), String> {
163        self.send_request(
164            host::SHELL_EXECUTE_COMMAND,
165            serde_json::json!({ "cmd": cmd }),
166        )
167        .await?;
168        Ok(())
169    }
170
171    pub async fn shell_open_folder(&self, path: &str) -> Result<(), String> {
172        self.send_request(host::SHELL_OPEN_FOLDER, serde_json::json!({ "path": path }))
173            .await?;
174        Ok(())
175    }
176
177    pub async fn shell_execute_elevation(&self, path: &str) -> Result<(), String> {
178        self.send_request(
179            host::SHELL_EXECUTE_ELEVATION,
180            serde_json::json!({ "path": path }),
181        )
182        .await?;
183        Ok(())
184    }
185
186    pub async fn notify(&self, title: &str, message: &str) -> Result<(), String> {
187        self.send_request(
188            host::NOTIFY,
189            serde_json::json!({ "title": title, "message": message }),
190        )
191        .await?;
192        Ok(())
193    }
194
195    /// 获取宿主当前界面语言(如 "zh-Hans"),供插件生成本地化文本。
196    pub async fn get_locale(&self) -> Result<String, String> {
197        let result = self
198            .send_request(host::GET_LOCALE, serde_json::json!(null))
199            .await?;
200        Ok(result.as_str().unwrap_or("").to_string())
201    }
202
203    /// 查询宿主当前实际生效主题,返回 `light` 或 `dark`。
204    pub async fn get_theme(&self) -> Result<String, String> {
205        let result = self
206            .send_request(host::GET_THEME, serde_json::Value::Null)
207            .await?;
208        result
209            .as_str()
210            .map(str::to_string)
211            .ok_or_else(|| "host theme response is not a string".to_string())
212    }
213
214    /// 全网模型清单(聚合所有提供方)。
215    pub async fn model_list(&self) -> Result<Vec<ModelInfo>, String> {
216        let result = self
217            .send_request(host::MODEL_LIST, serde_json::Value::Null)
218            .await?;
219        serde_json::from_value(result).map_err(|e| e.to_string())
220    }
221
222    /// 按 model_id 调用文本生成。
223    pub async fn model_chat(&self, req: ModelChatRequest) -> Result<ModelChatResponse, String> {
224        let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
225        let result = self.send_request(host::MODEL_CHAT, params).await?;
226        serde_json::from_value(result).map_err(|e| e.to_string())
227    }
228
229    /// 按 model_id 调用文本向量化(task_type 必填,宿主对缺失/未知值返回错误)。
230    pub async fn model_embedding(
231        &self,
232        req: ModelEmbeddingRequest,
233    ) -> Result<ModelEmbeddingResponse, String> {
234        let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
235        let result = self.send_request(host::MODEL_EMBEDDING, params).await?;
236        serde_json::from_value(result).map_err(|e| e.to_string())
237    }
238
239    /// 按 model_id 计算查询向量与多个目标向量的相似度。
240    pub async fn model_similarity(
241        &self,
242        req: ModelSimilarityRequest,
243    ) -> Result<ModelSimilarityResponse, String> {
244        let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
245        let result = self.send_request(host::MODEL_SIMILARITY, params).await?;
246        serde_json::from_value(result).map_err(|e| e.to_string())
247    }
248
249    pub async fn enumerate_apps(&self) -> Result<serde_json::Value, String> {
250        self.send_request(host::APP_ENUMERATE, serde_json::json!(null))
251            .await
252    }
253
254    pub async fn resolve_path(&self, kind: &str) -> Result<String, String> {
255        let result = self
256            .send_request(host::PATH_RESOLVE, serde_json::json!({ "kind": kind }))
257            .await?;
258        Ok(result.as_str().unwrap_or("").to_string())
259    }
260
261    /// 上传插件本地文件到宿主资源空间。
262    /// 直接传递文件路径,由宿主负责读取。
263    pub async fn resource_upload(
264        &self,
265        resource_id: &str,
266        file_path: &str,
267        max_size: Option<u64>,
268    ) -> Result<String, String> {
269        let result = self
270            .send_request(
271                host::RESOURCE_UPLOAD,
272                serde_json::json!({
273                    "resourceId": resource_id,
274                    "filePath": file_path,
275                    "maxSize": max_size,
276                }),
277            )
278            .await?;
279        Ok(result.as_str().unwrap_or("").to_string())
280    }
281
282    pub async fn resource_get(&self, resource_id: &str) -> Result<Vec<u8>, String> {
283        let result = self
284            .send_request(
285                host::RESOURCE_GET,
286                serde_json::json!({
287                    "resourceId": resource_id,
288                }),
289            )
290            .await?;
291        let b64 = result.as_str().unwrap_or("");
292        base64::engine::general_purpose::STANDARD
293            .decode(b64)
294            .map_err(|e| format!("base64 decode failed: {}", e))
295    }
296
297    /// 直接写入资源字节数据(无需临时文件),base64 编解码由 SDK 内部处理。
298    pub async fn resource_put(&self, resource_id: &str, data: &[u8]) -> Result<(), String> {
299        let b64 = base64::engine::general_purpose::STANDARD.encode(data);
300        self.send_request(
301            host::RESOURCE_PUT,
302            serde_json::json!({
303                "resourceId": resource_id,
304                "bytesB64": b64,
305            }),
306        )
307        .await?;
308        Ok(())
309    }
310
311    /// 删除资源文件。
312    pub async fn resource_delete(&self, resource_id: &str) -> Result<(), String> {
313        self.send_request(
314            host::RESOURCE_DELETE,
315            serde_json::json!({
316                "resourceId": resource_id,
317            }),
318        )
319        .await?;
320        Ok(())
321    }
322
323    /// 列出本插件的所有资源标识符。
324    pub async fn resource_list(&self) -> Result<Vec<String>, String> {
325        let result = self
326            .send_request(host::RESOURCE_LIST, serde_json::json!({}))
327            .await?;
328        serde_json::from_value(result).map_err(|e| format!("parse resource list failed: {}", e))
329    }
330}
331
332#[cfg(test)]
333mod tests {
334    use super::*;
335
336    fn make_proxy() -> (Arc<HostProxy>, mpsc::Receiver<Vec<u8>>) {
337        let (tx, rx) = mpsc::channel::<Vec<u8>>(16);
338        let pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>> =
339            Arc::new(DashMap::new());
340        (Arc::new(HostProxy::new(pending, tx)), rx)
341    }
342
343    /// 回归测试:host/* 请求必须以**干净 JSON** 投递到 outbound 通道,
344    /// 分帧由 write_task 统一完成。
345    ///
346    /// 修复前 send_request/log_no_wait 在投递前先 encode_frame,
347    /// 导致双重分帧(帧内嵌帧):写出的字节为
348    /// `Content-Length: M\r\n\r\nContent-Length: N\r\n\r\n{JSON}`,
349    /// 宿主 read_frame 后 body 以 `Content-Length:` 开头,serde_json 解析失败
350    /// (`expected value at line 1 column 1`),host/* 调用全部超时。
351    #[tokio::test]
352    async fn send_request_posts_clean_json_without_embedded_frame() {
353        let (proxy, mut rx) = make_proxy();
354
355        // log_no_wait 是同步 fire-and-forget:直接检查通道里的字节。
356        proxy.log_no_wait("warn", "test message");
357        let bytes = rx.recv().await.expect("log_no_wait 应投递一条消息");
358        let text = String::from_utf8(bytes.clone()).expect("UTF-8");
359        assert!(
360            !text.contains("Content-Length"),
361            "outbound 通道中的消息不得包含帧头(双重分帧): {:?}",
362            text
363        );
364        // 必须是合法 JSON-RPC 消息
365        let v: serde_json::Value = serde_json::from_slice(&bytes).expect("干净 JSON 可解析");
366        assert_eq!(v["method"], host::LOG);
367        assert_eq!(v["params"]["message"], "test message");
368    }
369
370    /// send_request 投递的也必须是干净 JSON(await 路径)。
371    #[tokio::test]
372    async fn send_request_await_path_posts_clean_json() {
373        let (proxy, mut rx) = make_proxy();
374        // send_request 会 await 响应(oneshot 永不完成→挂起),
375        // 因此只在通道上取一条消息验证负载,然后丢弃任务。
376        let task = tokio::spawn(async move {
377            let _ = proxy.model_list().await; // 挂起直到超时(30s),任务随测试结束
378        });
379        let bytes = tokio::time::timeout(Duration::from_secs(2), rx.recv())
380            .await
381            .expect("send_request 应投递一条消息")
382            .expect("通道未关闭");
383        let text = String::from_utf8_lossy(&bytes);
384        assert!(
385            !text.contains("Content-Length"),
386            "outbound 通道中的消息不得包含帧头(双重分帧): {:?}",
387            text
388        );
389        let v: serde_json::Value = serde_json::from_slice(&bytes).expect("干净 JSON 可解析");
390        assert_eq!(v["method"], host::MODEL_LIST);
391        // 清理:中止挂起的任务
392        task.abort();
393    }
394}