1use gonzalo_graph::{CodeGraph, Language};
16use serde::{Deserialize, Serialize};
17use std::path::{Path, PathBuf};
18use std::process::Stdio;
19use std::sync::atomic::{AtomicUsize, Ordering};
20use std::time::Duration;
21use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
22use tokio::process::{ChildStdin, ChildStdout, Command};
23use tokio::sync::Mutex;
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
29pub struct ParseRequest {
30 pub language: Language,
31 pub source: String,
32}
33
34#[derive(Debug, thiserror::Error)]
37pub enum ParseError {
38 #[error("failed to spawn parse worker: {0}")]
39 Spawn(#[source] std::io::Error),
40 #[error("parse worker died (crashed or closed its pipe)")]
41 WorkerDied,
42 #[error("parse worker exceeded the {0:?} timeout")]
43 Timeout(Duration),
44 #[error("parse worker sent a malformed response: {0}")]
45 Protocol(String),
46}
47
48pub struct ParserPool {
50 worker_bin: PathBuf,
51 worker_env: Vec<(String, String)>,
52 slots: Vec<Mutex<Option<Worker>>>,
53 next: AtomicUsize,
54 timeout: Duration,
55}
56
57impl ParserPool {
58 pub fn new(worker_bin: impl Into<PathBuf>, size: usize, timeout: Duration) -> Self {
62 let size = size.max(1);
63 let slots = (0..size).map(|_| Mutex::new(None)).collect();
64 Self {
65 worker_bin: worker_bin.into(),
66 worker_env: Vec::new(),
67 slots,
68 next: AtomicUsize::new(0),
69 timeout,
70 }
71 }
72
73 pub fn with_worker_env(mut self, vars: Vec<(String, String)>) -> Self {
76 self.worker_env = vars;
77 self
78 }
79
80 pub fn size(&self) -> usize {
82 self.slots.len()
83 }
84
85 pub async fn parse(&self, language: Language, source: &str) -> Result<CodeGraph, ParseError> {
89 let idx = self.next.fetch_add(1, Ordering::Relaxed) % self.slots.len();
90 let mut slot = self.slots[idx].lock().await;
91
92 let mut last_err = ParseError::WorkerDied;
93 for _ in 0..2 {
94 if slot.is_none() {
95 *slot = Some(
96 Worker::spawn(&self.worker_bin, &self.worker_env).map_err(ParseError::Spawn)?,
97 );
98 }
99 let worker = slot.as_mut().expect("worker present");
100 match tokio::time::timeout(self.timeout, worker.roundtrip(language, source)).await {
101 Ok(Ok(graph)) => return Ok(graph),
102 Ok(Err(e)) => {
103 *slot = None;
105 last_err = e;
106 }
107 Err(_) => {
108 *slot = None;
110 return Err(ParseError::Timeout(self.timeout));
111 }
112 }
113 }
114 Err(last_err)
115 }
116}
117
118struct Worker {
120 _child: tokio::process::Child,
122 stdin: ChildStdin,
123 stdout: BufReader<ChildStdout>,
124}
125
126impl Worker {
127 fn spawn(bin: &Path, env: &[(String, String)]) -> std::io::Result<Self> {
128 let mut child = Command::new(bin)
129 .envs(env.iter().map(|(k, v)| (k.as_str(), v.as_str())))
130 .stdin(Stdio::piped())
131 .stdout(Stdio::piped())
132 .stderr(Stdio::null())
133 .kill_on_drop(true)
134 .spawn()?;
135 let stdin = child.stdin.take().expect("stdin piped");
136 let stdout = BufReader::new(child.stdout.take().expect("stdout piped"));
137 Ok(Self {
138 _child: child,
139 stdin,
140 stdout,
141 })
142 }
143
144 async fn roundtrip(
147 &mut self,
148 language: Language,
149 source: &str,
150 ) -> Result<CodeGraph, ParseError> {
151 let request = ParseRequest {
152 language,
153 source: source.to_string(),
154 };
155 let mut req = serde_json::to_string(&request).expect("ParseRequest serializes");
156 req.push('\n');
157 self.stdin
158 .write_all(req.as_bytes())
159 .await
160 .map_err(|_| ParseError::WorkerDied)?;
161 self.stdin
162 .flush()
163 .await
164 .map_err(|_| ParseError::WorkerDied)?;
165
166 let mut line = String::new();
167 let n = self
168 .stdout
169 .read_line(&mut line)
170 .await
171 .map_err(|_| ParseError::WorkerDied)?;
172 if n == 0 {
173 return Err(ParseError::WorkerDied); }
175 serde_json::from_str(&line).map_err(|e| ParseError::Protocol(e.to_string()))
176 }
177}