use std::{future::Future, pin::Pin, str::FromStr, sync::Arc};
use bytes::Bytes;
use connectrpc::{ConnectRpcService, dispatcher::Dispatcher};
use tauri::{
Manager, Runtime,
ipc::{InvokeResponseBody, JavaScriptChannelId},
plugin::TauriPlugin,
};
use crate::{call, codec, registry::CallRegistry, scheme, wire};
pub const PLUGIN_NAME: &str = "connectrpc-tauri";
type StartFn = Box<
dyn Fn(
wire::StartRequest,
Option<tauri::ipc::Channel<InvokeResponseBody>>,
Arc<CallRegistry>,
) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, String>> + Send>>
+ Send
+ Sync,
>;
pub(crate) struct TransportState {
pub(crate) calls: Arc<CallRegistry>,
start: StartFn,
}
pub fn serve<R: Runtime, D: Dispatcher>(service: ConnectRpcService<D>) -> TauriPlugin<R> {
let builder = tauri::plugin::Builder::new(PLUGIN_NAME);
let builder = scheme::register(builder, service.clone());
builder
.setup(move |app, _api| {
let state = TransportState {
calls: Arc::new(CallRegistry::default()),
start: Box::new(move |start, channel, calls| {
let service = service.clone();
Box::pin(async move { call::start(service, calls, start, channel).await })
}),
};
app.manage(state);
Ok(())
})
.invoke_handler(tauri::generate_handler![
connect_rpc,
connect_rpc_send,
connect_rpc_cancel
])
.build()
}
#[tauri::command]
async fn connect_rpc<R: Runtime>(
webview: tauri::Webview<R>,
request: tauri::ipc::Request<'_>,
) -> Result<tauri::ipc::Response, String> {
let start: wire::StartRequest = codec::decode(request.body())?;
let channel = if start.channel.is_empty() {
None
} else {
Some(
JavaScriptChannelId::from_str(&start.channel)
.map_err(|e| format!("invalid response channel: {e}"))?
.channel_on(webview.clone()),
)
};
let state = webview.state::<TransportState>();
let calls = Arc::clone(&state.calls);
let bytes = (state.start)(start, channel, calls).await?;
Ok(tauri::ipc::Response::new(bytes))
}
#[tauri::command]
async fn connect_rpc_send<R: Runtime>(
webview: tauri::Webview<R>,
request: tauri::ipc::Request<'_>,
) -> Result<(), String> {
let send: wire::SendRequest = codec::decode(request.body())?;
let calls = Arc::clone(&webview.state::<TransportState>().calls);
if send.end_of_stream {
calls.close_request_body(send.call_id);
return Ok(());
}
let Some(tx) = calls.body_sender(send.call_id) else {
return Ok(());
};
tx.send(Bytes::from(send.chunk))
.await
.map_err(|_| "request body closed".to_string())
}
#[tauri::command]
async fn connect_rpc_cancel<R: Runtime>(
webview: tauri::Webview<R>,
request: tauri::ipc::Request<'_>,
) -> Result<(), String> {
let cancel: wire::CancelRequest = codec::decode(request.body())?;
webview
.state::<TransportState>()
.calls
.remove(cancel.call_id);
Ok(())
}