use dashmap::DashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{mpsc, oneshot};
use zerolaunch_plugin_api::services::model::{
ModelChatRequest, ModelChatResponse, ModelEmbeddingRequest, ModelEmbeddingResponse, ModelInfo,
ModelSimilarityRequest, ModelSimilarityResponse,
};
use zerolaunch_plugin_protocol::methods::host;
use zerolaunch_plugin_protocol::JsonRpcError;
use base64::Engine as _;
pub struct HostProxy {
next_id: AtomicU64,
pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>>,
outbound_tx: mpsc::Sender<Vec<u8>>,
}
impl HostProxy {
pub fn new(
pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>>,
outbound_tx: mpsc::Sender<Vec<u8>>,
) -> Self {
Self {
next_id: AtomicU64::new(1),
pending,
outbound_tx,
}
}
async fn send_request(
&self,
method: &str,
params: serde_json::Value,
) -> Result<serde_json::Value, String> {
let is_model_call = method.starts_with("host/model.");
let timeout = if is_model_call {
Duration::from_secs(300)
} else {
Duration::from_secs(30)
};
self.send_request_inner(method, params, timeout).await
}
async fn send_request_inner(
&self,
method: &str,
params: serde_json::Value,
timeout: Duration,
) -> Result<serde_json::Value, String> {
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let request = serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
});
let payload = serde_json::to_vec(&request).map_err(|e| e.to_string())?;
let (tx, rx) = oneshot::channel();
self.pending.insert(id, tx);
if self.outbound_tx.send(payload).await.is_err() {
self.pending.remove(&id);
return Err("write channel closed".to_string());
}
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(Ok(value))) => Ok(value),
Ok(Ok(Err(err))) => Err(format!(
"host call failed (code {}): {}",
err.code, err.message
)),
Ok(Err(_)) => Err("response channel closed".to_string()),
Err(_) => {
self.pending.remove(&id);
Err("host call timed out".to_string())
}
}
}
pub async fn log(&self, level: &str, message: &str) -> Result<(), String> {
self.send_request(
host::LOG,
serde_json::json!({ "level": level, "message": message }),
)
.await?;
Ok(())
}
pub fn log_no_wait(&self, level: &str, message: &str) {
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let Ok(payload) = serde_json::to_vec(&serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": host::LOG,
"params": { "level": level, "message": message },
})) else {
return;
};
let (tx, _rx) = oneshot::channel(); self.pending.insert(id, tx);
if self.outbound_tx.try_send(payload).is_err() {
self.pending.remove(&id);
}
}
pub async fn shell_open(&self, target: &str) -> Result<(), String> {
self.send_request(host::SHELL_OPEN, serde_json::json!({ "target": target }))
.await?;
Ok(())
}
pub async fn get_icon(&self, path: &str) -> Result<String, String> {
let result = self
.send_request(
host::ICON_GET,
serde_json::json!({ "request": { "path": path }, "level": "Full" }),
)
.await?;
Ok(result.as_str().unwrap_or("").to_string())
}
pub async fn shell_execute_command(&self, cmd: &str) -> Result<(), String> {
self.send_request(
host::SHELL_EXECUTE_COMMAND,
serde_json::json!({ "cmd": cmd }),
)
.await?;
Ok(())
}
pub async fn shell_open_folder(&self, path: &str) -> Result<(), String> {
self.send_request(host::SHELL_OPEN_FOLDER, serde_json::json!({ "path": path }))
.await?;
Ok(())
}
pub async fn shell_execute_elevation(&self, path: &str) -> Result<(), String> {
self.send_request(
host::SHELL_EXECUTE_ELEVATION,
serde_json::json!({ "path": path }),
)
.await?;
Ok(())
}
pub async fn notify(&self, title: &str, message: &str) -> Result<(), String> {
self.send_request(
host::NOTIFY,
serde_json::json!({ "title": title, "message": message }),
)
.await?;
Ok(())
}
pub async fn get_locale(&self) -> Result<String, String> {
let result = self
.send_request(host::GET_LOCALE, serde_json::json!(null))
.await?;
Ok(result.as_str().unwrap_or("").to_string())
}
pub async fn get_theme(&self) -> Result<String, String> {
let result = self
.send_request(host::GET_THEME, serde_json::Value::Null)
.await?;
result
.as_str()
.map(str::to_string)
.ok_or_else(|| "host theme response is not a string".to_string())
}
pub async fn model_list(&self) -> Result<Vec<ModelInfo>, String> {
let result = self
.send_request(host::MODEL_LIST, serde_json::Value::Null)
.await?;
serde_json::from_value(result).map_err(|e| e.to_string())
}
pub async fn model_chat(&self, req: ModelChatRequest) -> Result<ModelChatResponse, String> {
let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
let result = self.send_request(host::MODEL_CHAT, params).await?;
serde_json::from_value(result).map_err(|e| e.to_string())
}
pub async fn model_embedding(
&self,
req: ModelEmbeddingRequest,
) -> Result<ModelEmbeddingResponse, String> {
let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
let result = self.send_request(host::MODEL_EMBEDDING, params).await?;
serde_json::from_value(result).map_err(|e| e.to_string())
}
pub async fn model_similarity(
&self,
req: ModelSimilarityRequest,
) -> Result<ModelSimilarityResponse, String> {
let params = serde_json::to_value(req).map_err(|e| e.to_string())?;
let result = self.send_request(host::MODEL_SIMILARITY, params).await?;
serde_json::from_value(result).map_err(|e| e.to_string())
}
pub async fn enumerate_apps(&self) -> Result<serde_json::Value, String> {
self.send_request(host::APP_ENUMERATE, serde_json::json!(null))
.await
}
pub async fn resolve_path(&self, kind: &str) -> Result<String, String> {
let result = self
.send_request(host::PATH_RESOLVE, serde_json::json!({ "kind": kind }))
.await?;
Ok(result.as_str().unwrap_or("").to_string())
}
pub async fn resource_upload(
&self,
resource_id: &str,
file_path: &str,
max_size: Option<u64>,
) -> Result<String, String> {
let result = self
.send_request(
host::RESOURCE_UPLOAD,
serde_json::json!({
"resourceId": resource_id,
"filePath": file_path,
"maxSize": max_size,
}),
)
.await?;
Ok(result.as_str().unwrap_or("").to_string())
}
pub async fn resource_get(&self, resource_id: &str) -> Result<Vec<u8>, String> {
let result = self
.send_request(
host::RESOURCE_GET,
serde_json::json!({
"resourceId": resource_id,
}),
)
.await?;
let b64 = result.as_str().unwrap_or("");
base64::engine::general_purpose::STANDARD
.decode(b64)
.map_err(|e| format!("base64 decode failed: {}", e))
}
pub async fn resource_put(&self, resource_id: &str, data: &[u8]) -> Result<(), String> {
let b64 = base64::engine::general_purpose::STANDARD.encode(data);
self.send_request(
host::RESOURCE_PUT,
serde_json::json!({
"resourceId": resource_id,
"bytesB64": b64,
}),
)
.await?;
Ok(())
}
pub async fn resource_delete(&self, resource_id: &str) -> Result<(), String> {
self.send_request(
host::RESOURCE_DELETE,
serde_json::json!({
"resourceId": resource_id,
}),
)
.await?;
Ok(())
}
pub async fn resource_list(&self) -> Result<Vec<String>, String> {
let result = self
.send_request(host::RESOURCE_LIST, serde_json::json!({}))
.await?;
serde_json::from_value(result).map_err(|e| format!("parse resource list failed: {}", e))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_proxy() -> (Arc<HostProxy>, mpsc::Receiver<Vec<u8>>) {
let (tx, rx) = mpsc::channel::<Vec<u8>>(16);
let pending: Arc<DashMap<u64, oneshot::Sender<Result<serde_json::Value, JsonRpcError>>>> =
Arc::new(DashMap::new());
(Arc::new(HostProxy::new(pending, tx)), rx)
}
#[tokio::test]
async fn send_request_posts_clean_json_without_embedded_frame() {
let (proxy, mut rx) = make_proxy();
proxy.log_no_wait("warn", "test message");
let bytes = rx.recv().await.expect("log_no_wait 应投递一条消息");
let text = String::from_utf8(bytes.clone()).expect("UTF-8");
assert!(
!text.contains("Content-Length"),
"outbound 通道中的消息不得包含帧头(双重分帧): {:?}",
text
);
let v: serde_json::Value = serde_json::from_slice(&bytes).expect("干净 JSON 可解析");
assert_eq!(v["method"], host::LOG);
assert_eq!(v["params"]["message"], "test message");
}
#[tokio::test]
async fn send_request_await_path_posts_clean_json() {
let (proxy, mut rx) = make_proxy();
let task = tokio::spawn(async move {
let _ = proxy.model_list().await; });
let bytes = tokio::time::timeout(Duration::from_secs(2), rx.recv())
.await
.expect("send_request 应投递一条消息")
.expect("通道未关闭");
let text = String::from_utf8_lossy(&bytes);
assert!(
!text.contains("Content-Length"),
"outbound 通道中的消息不得包含帧头(双重分帧): {:?}",
text
);
let v: serde_json::Value = serde_json::from_slice(&bytes).expect("干净 JSON 可解析");
assert_eq!(v["method"], host::MODEL_LIST);
task.abort();
}
}