1use crate::{
2 config::{Backend, Profile},
3 error::Error,
4};
5use base64::{Engine, engine::general_purpose::STANDARD};
6use serde_json::Value;
7use std::{path::PathBuf, process::Stdio, time::Duration};
8use tokio::{
9 io::{AsyncRead, AsyncReadExt, AsyncWriteExt},
10 process::Command,
11};
12
13pub const WRITE_SCRIPT: &str = include_str!("desktop_write.ps1");
14pub const SCRIPT: &str = include_str!("desktop.ps1");
15const MAX_OUTPUT: u64 = 32 * 1024 * 1024;
16
17pub fn powershell() -> Result<PathBuf, Error> {
18 let supported = cfg!(windows)
19 || (cfg!(target_os = "linux")
20 && std::fs::read_to_string("/proc/sys/kernel/osrelease")
21 .unwrap_or_default()
22 .to_ascii_lowercase()
23 .contains("microsoft"));
24 if !supported {
25 return Err(Error::new(
26 "unsupported",
27 "Desktop access requires Windows or WSL with Windows interop, and the Windows OneNote desktop application. macOS and web OneNote do not expose this COM interface.",
28 ));
29 }
30 std::env::split_paths(&std::env::var_os("PATH").unwrap_or_default())
31 .map(|p| p.join("powershell.exe")).find(|p| p.is_file())
32 .ok_or_else(|| Error::desktop("powershell.exe is not on PATH; enable Windows interop in WSL or restore Windows PowerShell to PATH"))
33}
34
35pub const REMOTE_SCRIPT: &str = include_str!("desktop_remote.ps1");
36
37pub fn executable(profile: &Profile) -> Result<PathBuf, Error> {
38 if profile.backend == Backend::Desktop {
39 return powershell();
40 }
41 let name = if cfg!(windows) { "ssh.exe" } else { "ssh" };
42 std::env::split_paths(&std::env::var_os("PATH").unwrap_or_default())
43 .map(|p| p.join(name))
44 .find(|p| p.is_file())
45 .ok_or_else(|| {
46 Error::desktop("OpenSSH client is not on PATH; install it before using an SSH profile")
47 })
48}
49fn ssh_command(profile: &Profile) -> Result<Command, Error> {
50 profile.validate()?;
51 let mut command = Command::new(executable(profile)?);
52 command.args([
53 "-T",
54 "-o",
55 "BatchMode=yes",
56 "-o",
57 "ConnectTimeout=10",
58 "-o",
59 "ServerAliveInterval=10",
60 "-o",
61 "ServerAliveCountMax=2",
62 "-o",
63 "StrictHostKeyChecking=yes",
64 ]);
65 if let Some(path) = &profile.identity_file {
66 command.arg("-i").arg(path);
67 }
68 if let Some(port) = profile.port {
69 command.arg("-p").arg(port.to_string());
70 }
71 command
72 .arg("--")
73 .arg(profile.host.as_deref().unwrap_or_default())
74 .arg("powershell.exe");
75 Ok(command)
76}
77pub async fn run(profile: &Profile, request: Value) -> Result<String, Error> {
78 let value = scan(profile, request).await?;
79 value["xml"]
80 .as_str()
81 .map(str::to_owned)
82 .ok_or_else(|| Error::desktop("OneNote bridge returned no XML"))
83}
84pub async fn scan(profile: &Profile, request: Value) -> Result<Value, Error> {
85 execute(profile, request, SCRIPT).await
86}
87pub async fn write(profile: &Profile, request: Value) -> Result<Value, Error> {
88 if profile.read_only {
89 return Err(Error::new(
90 "read_only",
91 "This profile is read-only; select a writable profile to create or append pages",
92 ));
93 }
94 execute(profile, request, WRITE_SCRIPT).await
95}
96async fn execute(profile: &Profile, request: Value, script: &str) -> Result<Value, Error> {
97 if profile.backend == Backend::Desktop {
98 run_bridge(
99 Command::new(powershell()?),
100 request,
101 script,
102 Duration::from_secs(45),
103 )
104 .await
105 } else {
106 run_bridge_inner(
107 ssh_command(profile)?,
108 request,
109 script,
110 Duration::from_secs(90),
111 true,
112 )
113 .await.map_err(|mut error| {
114 if error.message.contains("Invalid PowerShell response") {
115 error.message.push_str(" Check SSH key/agent authentication, verify the host key with your SSH client, and confirm Windows PowerShell is available on the destination.");
116 }
117 error
118 })
119 }
120}
121
122async fn bounded(reader: impl AsyncRead + Unpin, max: u64) -> Result<Vec<u8>, Error> {
123 let mut bytes = Vec::new();
124 reader.take(max + 1).read_to_end(&mut bytes).await?;
125 if bytes.len() as u64 > max {
126 return Err(Error::desktop(
127 "OneNote response exceeded the size limit; narrow the notebook or search scope",
128 ));
129 }
130 Ok(bytes)
131}
132
133async fn run_bridge(
134 command: Command,
135 request: Value,
136 script: &str,
137 timeout: Duration,
138) -> Result<Value, Error> {
139 run_bridge_inner(command, request, script, timeout, false).await
140}
141async fn run_bridge_inner(
142 mut command: Command,
143 request: Value,
144 script: &str,
145 timeout: Duration,
146 remote: bool,
147) -> Result<Value, Error> {
148 let executable_script = if remote { REMOTE_SCRIPT } else { script };
149 let encoded = STANDARD.encode(
150 executable_script
151 .encode_utf16()
152 .flat_map(u16::to_le_bytes)
153 .collect::<Vec<_>>(),
154 );
155 let mut child = command
156 .args([
157 "-NoLogo",
158 "-NoProfile",
159 "-NonInteractive",
160 "-STA",
161 "-EncodedCommand",
162 &encoded,
163 ])
164 .stdin(Stdio::piped())
165 .stdout(Stdio::piped())
166 .stderr(Stdio::piped())
167 .kill_on_drop(true)
168 .spawn()
169 .map_err(|e| {
170 Error::desktop(format!(
171 "Cannot start {}: {e}",
172 if remote { "SSH" } else { "Windows PowerShell" }
173 ))
174 })?;
175 let mut stdin = child
176 .stdin
177 .take()
178 .ok_or_else(|| Error::desktop("Missing bridge stdin"))?;
179 let stdout = child
180 .stdout
181 .take()
182 .ok_or_else(|| Error::desktop("Missing bridge stdout"))?;
183 let stderr = child
184 .stderr
185 .take()
186 .ok_or_else(|| Error::desktop("Missing bridge stderr"))?;
187 let payload = if remote {
188 serde_json::json!({"script_base64":STANDARD.encode(format!("\u{feff}{script}")),"request":request,"write":script == WRITE_SCRIPT})
189 } else {
190 request.clone()
191 };
192 let bytes = serde_json::to_vec(&payload).map_err(|e| Error::invalid(e.to_string()))?;
193 let uncertain = |e: Error| {
194 if script == WRITE_SCRIPT {
195 Error::new(
196 "write_uncertain",
197 format!(
198 "{}. The {} on target {} may have taken effect; inspect OneNote before retrying.",
199 e.message,
200 request["operation"].as_str().unwrap_or("write"),
201 request["target"].as_str().unwrap_or("unknown")
202 ),
203 )
204 } else {
205 e
206 }
207 };
208 let operation = async {
209 let (_, stdout, stderr, status) = tokio::try_join!(
210 async {
211 stdin.write_all(&bytes).await?;
212 drop(stdin);
213 Ok::<_, Error>(())
214 },
215 bounded(stdout, MAX_OUTPUT),
216 bounded(stderr, 64 * 1024),
217 async { Ok::<_, Error>(child.wait().await?) },
218 )
219 .map_err(uncertain)?;
220 if script == WRITE_SCRIPT {
222 let decoded: Option<Value> =
223 serde_json::from_slice(stdout.strip_prefix(&[0xef, 0xbb, 0xbf]).unwrap_or(&stdout))
224 .ok();
225 if decoded.as_ref().is_some_and(|v| v.get("error").is_some()) {
226 return parse_output(status.success(), &stdout, &stderr);
227 }
228 parse_output(status.success(), &stdout, &stderr).map_err(uncertain)
229 } else {
230 parse_output(status.success(), &stdout, &stderr)
231 }
232 };
233 let result = match tokio::time::timeout(timeout, operation).await {
234 Ok(result) => result,
235 Err(_) => Err(uncertain(Error::desktop(format!(
236 "OneNote did not respond within {} seconds; check its Windows desktop for setup, password, or sync prompts",
237 timeout.as_secs()
238 )))),
239 };
240 if result.is_err() {
241 let _ = child.kill().await;
242 }
243 result
244}
245
246fn parse_output(success: bool, stdout: &[u8], stderr: &[u8]) -> Result<Value, Error> {
247 let value: Value =
248 serde_json::from_slice(stdout.strip_prefix(&[0xef, 0xbb, 0xbf]).unwrap_or(stdout))
249 .map_err(|_| {
250 Error::desktop(format!(
251 "Invalid PowerShell response: {}",
252 String::from_utf8_lossy(stderr).trim()
253 ))
254 })?;
255 if let Some(error) = value.get("error") {
256 let kind = match error["kind"].as_str() {
257 Some("invalid_input") => "invalid_input",
258 Some("not_found") => "not_found",
259 Some("conflict") => "conflict",
260 Some("write_uncertain") => "write_uncertain",
261 _ => "desktop_error",
262 };
263 let mut failure = Error::new(
264 kind,
265 error["message"]
266 .as_str()
267 .unwrap_or("OneNote COM operation failed"),
268 );
269 failure.page_id = error["page_id"]
270 .as_str()
271 .filter(|id| !id.is_empty())
272 .map(str::to_owned);
273 return Err(failure);
274 }
275 if !success {
276 return Err(Error::desktop(format!(
277 "PowerShell failed: {}",
278 String::from_utf8_lossy(stderr).trim()
279 )));
280 }
281 Ok(value)
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287 #[test]
288 fn script_fits_windows_command_line_and_protocol_errors_are_preserved() {
289 for script in [SCRIPT, WRITE_SCRIPT, REMOTE_SCRIPT] {
290 assert!(script.encode_utf16().count() * 2 * 4 / 3 < 32000);
291 }
292 assert!(parse_output(true, b"noise", b"").is_err());
293 assert_eq!(
294 parse_output(
295 false,
296 br#"{"error":{"kind":"not_found","message":"Gone"}}"#,
297 b""
298 )
299 .unwrap_err()
300 .kind,
301 "not_found"
302 );
303 assert_eq!(
304 parse_output(true, b"\xef\xbb\xbf{\"xml\":\"hello\"}\r\n", b"").unwrap()["xml"],
305 "hello"
306 );
307 }
308 #[tokio::test]
309 async fn output_size_is_bounded() {
310 assert!(bounded(&b"12345"[..], 4).await.is_err());
311 assert_eq!(bounded(&b"1234"[..], 4).await.unwrap(), b"1234");
312 }
313 #[cfg(unix)]
314 #[tokio::test]
315 async fn hung_bridge_times_out() {
316 let mut command = Command::new("sh");
317 command.args(["-c", "exec sleep 10"]);
318 let error = run_bridge(
319 command,
320 serde_json::json!({}),
321 SCRIPT,
322 Duration::from_millis(50),
323 )
324 .await
325 .unwrap_err();
326 assert_eq!(error.kind, "desktop_error");
327 assert!(error.message.contains("did not respond"));
328 }
329 #[cfg(unix)]
330 #[tokio::test]
331 async fn write_transport_failures_are_uncertain_but_com_errors_are_preserved() {
332 for (shell, expected) in [
333 ("exec sleep 10", "write_uncertain"),
334 ("cat >/dev/null; printf noise", "write_uncertain"),
335 (
336 "cat >/dev/null; printf '%s' '{\"error\":{\"kind\":\"desktop_error\",\"message\":\"Locked\",\"page_id\":\"page\"}}'; exit 1",
337 "desktop_error",
338 ),
339 ] {
340 let mut command = Command::new("sh");
341 command.args(["-c", shell]);
342 let error = run_bridge(
343 command,
344 serde_json::json!({"operation":"append","target":"page"}),
345 WRITE_SCRIPT,
346 Duration::from_millis(100),
347 )
348 .await
349 .unwrap_err();
350 assert_eq!(error.kind, expected);
351 if expected == "desktop_error" {
352 assert_eq!(error.page_id.as_deref(), Some("page"));
353 } else {
354 assert!(error.message.contains("inspect OneNote before retrying"));
355 }
356 }
357 }
358}