1use crate::checker::{Diagnostic, Severity};
2use anyhow::Result;
3use serde::Deserialize;
4use std::collections::HashMap;
5use tracing::{debug, warn};
6
7use super::Engine;
8
9pub struct ProselintEngine {
10 config_path: Option<String>,
11}
12
13impl ProselintEngine {
14 #[must_use]
15 pub const fn new(config_path: Option<String>) -> Self {
16 Self { config_path }
17 }
18}
19
20#[derive(Deserialize)]
22struct ProselintOutput {
23 result: HashMap<String, ProselintFileResult>,
24}
25
26#[derive(Deserialize)]
28#[serde(untagged)]
29enum ProselintFileResult {
30 Ok {
31 diagnostics: Vec<ProselintDiagnostic>,
32 },
33 Err {
34 error: ProselintError,
35 },
36}
37
38#[derive(Deserialize)]
39struct ProselintDiagnostic {
40 check_path: String,
41 message: String,
42 span: (usize, usize),
44 replacements: Option<String>,
46}
47
48#[derive(Deserialize)]
49struct ProselintError {
50 message: String,
51}
52
53#[allow(clippy::cast_possible_truncation)]
60fn char_span_to_byte_range(text: &str, span: (usize, usize)) -> (u32, u32) {
61 let char_start = span.0.saturating_sub(1);
63 let char_end = span.1.saturating_sub(1);
64
65 let mut byte_start = text.len();
66 let mut byte_end = text.len();
67
68 for (i, (byte_idx, _)) in text.char_indices().enumerate() {
69 if i == char_start {
70 byte_start = byte_idx;
71 }
72 if i == char_end {
73 byte_end = byte_idx;
74 break;
75 }
76 }
77
78 (byte_start as u32, byte_end as u32)
79}
80
81#[async_trait::async_trait]
82impl Engine for ProselintEngine {
83 fn name(&self) -> &'static str {
84 "proselint"
85 }
86
87 fn supported_languages(&self) -> Vec<String> {
88 vec!["en".to_string()]
89 }
90
91 async fn check(&mut self, text: &str, _language_id: &str) -> Result<Vec<Diagnostic>> {
92 use tokio::io::AsyncWriteExt;
93 use tokio::process::Command;
94
95 let mut cmd = Command::new("proselint");
96 cmd.arg("check").arg("-o").arg("json");
97
98 if let Some(cfg) = &self.config_path {
99 cmd.arg("--config").arg(cfg);
100 }
101
102 cmd.stdin(std::process::Stdio::piped())
103 .stdout(std::process::Stdio::piped())
104 .stderr(std::process::Stdio::piped());
105
106 let output = match cmd.spawn() {
107 Ok(mut child) => {
108 if let Some(mut stdin) = child.stdin.take() {
109 let _ = stdin.write_all(text.as_bytes()).await;
110 let _ = stdin.shutdown().await;
111 }
112 child.wait_with_output().await?
113 }
114 Err(e) => {
115 warn!("Failed to spawn proselint: {e}");
116 return Ok(vec![]);
117 }
118 };
119
120 let code = output.status.code().unwrap_or(4);
123 if code >= 2 {
124 let stderr = String::from_utf8_lossy(&output.stderr);
125 warn!(code, stderr = stderr.trim(), "Proselint error");
126 return Ok(vec![]);
127 }
128
129 let stdout = String::from_utf8_lossy(&output.stdout);
130 if stdout.trim().is_empty() {
131 return Ok(vec![]);
132 }
133
134 let mut de = serde_json::Deserializer::from_str(&stdout).into_iter::<ProselintOutput>();
137 let parsed: ProselintOutput = match de.next() {
138 Some(Ok(o)) => o,
139 Some(Err(e)) => {
140 warn!("Failed to parse proselint JSON: {e}");
141 debug!(stdout = %stdout, "Raw proselint output");
142 return Ok(vec![]);
143 }
144 None => return Ok(vec![]),
145 };
146
147 let mut diagnostics = Vec::new();
148 for file_result in parsed.result.into_values() {
149 match file_result {
150 ProselintFileResult::Ok { diagnostics: diags } => {
151 for d in diags {
152 let (start_byte, end_byte) = char_span_to_byte_range(text, d.span);
153 let suggestions = d.replacements.map(|r| vec![r]).unwrap_or_default();
154
155 diagnostics.push(Diagnostic {
156 start_byte,
157 end_byte,
158 message: d.message,
159 suggestions,
160 rule_id: format!("proselint.{}", d.check_path),
161 severity: Severity::Warning as i32,
162 unified_id: String::new(),
163 confidence: 0.7,
164 language: String::new(),
165 pack_installable: false,
166 });
167 }
168 }
169 ProselintFileResult::Err { error } => {
170 warn!(msg = error.message, "Proselint reported a file error");
171 }
172 }
173 }
174
175 Ok(diagnostics)
176 }
177}
178
179#[cfg(test)]
180mod tests {
181 use super::*;
182
183 #[test]
184 fn char_span_basic() {
185 let text = "Hello world";
186 let (start, end) = char_span_to_byte_range(text, (7, 12));
188 assert_eq!(start, 6);
189 assert_eq!(end, 11);
190 assert_eq!(&text[start as usize..end as usize], "world");
191 }
192
193 #[test]
194 fn char_span_start_of_text() {
195 let text = "Hello";
196 let (start, end) = char_span_to_byte_range(text, (1, 6));
198 assert_eq!(start, 0);
199 assert_eq!(end, 5);
200 assert_eq!(&text[start as usize..end as usize], "Hello");
201 }
202
203 #[test]
204 fn char_span_unicode() {
205 let text = "café latte";
206 let (start, end) = char_span_to_byte_range(text, (6, 11));
208 assert_eq!(&text[start as usize..end as usize], "latte");
209 }
210
211 #[test]
212 fn char_span_clamped() {
213 let text = "short";
214 let (start, end) = char_span_to_byte_range(text, (1, 100));
215 assert_eq!(start, 0);
216 assert_eq!(end as usize, text.len());
217 }
218
219 #[test]
220 fn proselint_diagnostic_deserializes() {
221 let json = r#"{
222 "check_path": "uncomparables",
223 "message": "Comparison of an uncomparable: 'very unique'.",
224 "span": [10, 21],
225 "replacements": "unique",
226 "pos": [1, 9]
227 }"#;
228 let d: ProselintDiagnostic = serde_json::from_str(json).unwrap();
229 assert_eq!(d.check_path, "uncomparables");
230 assert_eq!(d.span, (10, 21));
231 assert_eq!(d.replacements.as_deref(), Some("unique"));
232 }
233
234 #[test]
235 fn proselint_diagnostic_null_replacements() {
236 let json = r#"{
237 "check_path": "hedging",
238 "message": "Hedging: 'I think'.",
239 "span": [1, 8],
240 "replacements": null,
241 "pos": [1, 0]
242 }"#;
243 let d: ProselintDiagnostic = serde_json::from_str(json).unwrap();
244 assert!(d.replacements.is_none());
245 }
246
247 #[test]
248 fn proselint_full_output_deserializes() {
249 let json = r#"{
250 "result": {
251 "<stdin>": {
252 "diagnostics": [
253 {
254 "check_path": "uncomparables",
255 "message": "Comparison of an uncomparable.",
256 "span": [10, 21],
257 "replacements": "unique",
258 "pos": [1, 9]
259 }
260 ]
261 }
262 }
263 }"#;
264 let output: ProselintOutput = serde_json::from_str(json).unwrap();
265 assert_eq!(output.result.len(), 1);
266 match &output.result["<stdin>"] {
267 ProselintFileResult::Ok { diagnostics } => {
268 assert_eq!(diagnostics.len(), 1);
269 assert_eq!(diagnostics[0].check_path, "uncomparables");
270 }
271 ProselintFileResult::Err { .. } => panic!("expected Ok"),
272 }
273 }
274
275 #[test]
276 fn proselint_error_result_deserializes() {
277 let json = r#"{
278 "result": {
279 "<stdin>": {
280 "error": {
281 "code": -31997,
282 "message": "Some error occurred"
283 }
284 }
285 }
286 }"#;
287 let output: ProselintOutput = serde_json::from_str(json).unwrap();
288 match &output.result["<stdin>"] {
289 ProselintFileResult::Err { error } => {
290 assert_eq!(error.message, "Some error occurred");
291 }
292 ProselintFileResult::Ok { .. } => panic!("expected Err"),
293 }
294 }
295
296 #[tokio::test]
297 async fn proselint_engine_missing_binary() -> Result<()> {
298 let mut engine = ProselintEngine::new(None);
299 let result = engine.check("test text", "en-US").await;
300 assert!(result.is_ok());
301 Ok(())
302 }
303
304 #[tokio::test]
307 #[ignore]
308 async fn proselint_engine_live() -> Result<()> {
309 let mut engine = ProselintEngine::new(None);
310 let text = "This is very unique and extremely obvious.";
311 let diagnostics = engine.check(text, "en-US").await?;
312
313 println!("Proselint returned {} diagnostics:", diagnostics.len());
314 for d in &diagnostics {
315 println!(
316 " [{}-{}] {} (rule: {}, suggestions: {:?})",
317 d.start_byte, d.end_byte, d.message, d.rule_id, d.suggestions
318 );
319 }
320
321 assert!(
322 !diagnostics.is_empty(),
323 "Expected at least 1 diagnostic from proselint"
324 );
325 assert!(diagnostics[0].rule_id.starts_with("proselint."));
326 Ok(())
327 }
328}