1use std::collections::BTreeMap;
12use std::num::NonZeroU64;
13use std::time::Duration;
14
15use async_trait::async_trait;
16use onetaskgraph_plugin_api::{
17 Capabilities, DependencyEdge, Direction, Document, DocumentQuery, Health, ItemWrite, Label,
18 Metering, NativeId, Page, PageRequest, Project, ProjectQuery, SourceError, SourceName, Task,
19 TaskQuery, TaskSource, WriteSupport,
20};
21use serde::Deserialize;
22use serde_json::{Value, json};
23
24use super::connection::{Connection, Peer};
25use super::wire::{
26 DeleteParams, DependencyParams, DocumentQueryParams, DocumentResult, DocumentWriteParams,
27 EngineIdentity, IdParams, InitializeParams, InitializeResult, LabelParams, MeteringResult,
28 PROTOCOL_VERSION, ProjectQueryParams, ProjectResult, ProjectWriteParams, Request,
29 TaskQueryParams, TaskResult, TaskWriteParams, WriteResult,
30};
31
32const HANDSHAKE_ID: &str = "0";
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub struct RequestDeadline(NonZeroU64);
40
41impl RequestDeadline {
42 pub const DEFAULT: Self = Self(NonZeroU64::new(30_000).expect("non-zero default"));
44
45 #[must_use]
47 pub const fn from_millis(milliseconds: NonZeroU64) -> Self {
48 Self(milliseconds)
49 }
50
51 #[must_use]
53 pub const fn milliseconds(self) -> NonZeroU64 {
54 self.0
55 }
56
57 fn duration(self) -> Duration {
58 Duration::from_millis(self.0.get())
59 }
60}
61
62pub struct SubprocessSource {
64 kind: &'static str,
72 capabilities: Capabilities,
74 writes: WriteSupport,
80 meters: bool,
85 connection: Connection,
87}
88
89impl std::fmt::Debug for SubprocessSource {
90 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
93 f.debug_struct("SubprocessSource")
94 .field("kind", &self.kind)
95 .finish_non_exhaustive()
96 }
97}
98
99impl SubprocessSource {
100 pub fn connect(
109 program: &str,
110 args: &[String],
111 name: &SourceName,
112 config: &Value,
113 secrets: BTreeMap<String, String>,
114 ) -> Result<Self, SourceError> {
115 Self::connect_with_deadline(
116 program,
117 args,
118 name,
119 config,
120 secrets,
121 RequestDeadline::DEFAULT,
122 )
123 }
124
125 pub fn connect_with_deadline(
127 program: &str,
128 args: &[String],
129 name: &SourceName,
130 config: &Value,
131 secrets: BTreeMap<String, String>,
132 deadline: RequestDeadline,
133 ) -> Result<Self, SourceError> {
134 Self::adopt(
135 Peer::spawn(program, args, deadline.duration())?,
136 name,
137 config,
138 secrets,
139 )
140 }
141
142 pub fn over(
157 to_plugin: impl std::io::Write + Send + 'static,
158 from_plugin: impl std::io::Read + Send + 'static,
159 name: &SourceName,
160 config: &Value,
161 secrets: BTreeMap<String, String>,
162 ) -> Result<Self, SourceError> {
163 Self::over_with_request_deadline(
164 to_plugin,
165 from_plugin,
166 name,
167 config,
168 secrets,
169 RequestDeadline::DEFAULT,
170 )
171 }
172
173 pub fn over_with_request_deadline(
179 to_plugin: impl std::io::Write + Send + 'static,
180 from_plugin: impl std::io::Read + Send + 'static,
181 name: &SourceName,
182 config: &Value,
183 secrets: BTreeMap<String, String>,
184 deadline: RequestDeadline,
185 ) -> Result<Self, SourceError> {
186 Self::adopt(
187 Peer::over(to_plugin, from_plugin, deadline.duration()),
188 name,
189 config,
190 secrets,
191 )
192 }
193
194 fn adopt(
196 mut peer: Peer,
197 name: &SourceName,
198 config: &Value,
199 secrets: BTreeMap<String, String>,
200 ) -> Result<Self, SourceError> {
201 let result = Self::handshake(&mut peer, name, config, secrets);
202 let InitializeResult {
203 protocol_version,
204 kind,
205 capabilities,
206 writes,
207 meters,
208 } = match result {
209 Ok(result) => result,
210 Err(error) => return Err(with_diagnostics(error, &mut peer)),
211 };
212 let kind = kind.into_string();
213 if protocol_version != Some(PROTOCOL_VERSION) {
214 return Err(SourceError::Config {
215 message: match protocol_version {
216 Some(spoken) => format!(
217 "the {kind:?} plugin was asked for protocol version \
218 {PROTOCOL_VERSION} and answered in version {spoken}; the two are \
219 incompatible and this engine does not guess between them"
220 ),
221 None => format!(
222 "the {kind:?} plugin did not say which protocol version it \
223 answered in; this engine speaks version {PROTOCOL_VERSION} and \
224 does not guess"
225 ),
226 },
227 });
228 }
229 Ok(Self {
230 kind: String::leak(kind),
231 capabilities,
232 writes: writes.unwrap_or(WriteSupport::Unsupported),
233 meters,
234 connection: Connection::adopt(peer),
235 })
236 }
237
238 fn handshake(
240 peer: &mut Peer,
241 name: &SourceName,
242 config: &Value,
243 secrets: BTreeMap<String, String>,
244 ) -> Result<InitializeResult, SourceError> {
245 let params = InitializeParams {
246 protocol_version: PROTOCOL_VERSION,
247 engine: EngineIdentity {
248 name: "onetaskgraph".to_owned(),
249 version: env!("CARGO_PKG_VERSION").to_owned(),
250 },
251 source_name: name.as_str().to_owned(),
252 config: config.clone(),
253 secrets,
254 };
255 let request = Request {
256 id: HANDSHAKE_ID.to_owned(),
257 method: "initialize".to_owned(),
258 params: serde_json::to_value(¶ms).expect("a handshake is plain data"),
261 };
262 let line = peer.exchange(
263 &serde_json::to_string(&request).expect("a handshake request is plain data"),
264 )?;
265 let response: super::wire::Response =
266 serde_json::from_str(&line).map_err(|error| SourceError::Malformed {
267 message: format!(
268 "the plugin's handshake answer is not a response envelope: {error}"
269 ),
270 })?;
271 if response.id != HANDSHAKE_ID {
276 return Err(SourceError::Malformed {
277 message: format!(
278 "the plugin answered the handshake with an envelope addressed to {:?} \
279 rather than to {HANDSHAKE_ID:?}",
280 response.id
281 ),
282 });
283 }
284 let outcome = response.outcome().ok_or_else(|| SourceError::Malformed {
285 message: "the plugin's handshake answer carried both a result and an error, or \
286 neither"
287 .to_owned(),
288 })?;
289 let result = outcome?;
290 serde_json::from_value(result).map_err(|error| SourceError::Malformed {
291 message: format!("the plugin's handshake answer is not an initialize result: {error}"),
292 })
293 }
294
295 async fn ask<T: for<'de> Deserialize<'de>>(
297 &self,
298 method: &str,
299 params: Value,
300 ) -> Result<T, SourceError> {
301 let result = self.connection.call(method, params).await?;
302 serde_json::from_value(result).map_err(|error| SourceError::Malformed {
303 message: format!(
304 "the plugin's answer to {method} is not the shape it promises: {error}"
305 ),
306 })
307 }
308}
309
310fn with_diagnostics(error: SourceError, peer: &mut Peer) -> SourceError {
315 let said = peer.said();
316 if said.is_empty() {
317 return error;
318 }
319 let message = format!("{error}; the plugin wrote: {said}");
320 match error {
321 SourceError::RateLimited {
324 retry_after_seconds,
325 ..
326 } => SourceError::RateLimited {
327 retry_after_seconds,
328 message: Some(message),
329 },
330 SourceError::Config { .. } => SourceError::Config { message },
331 SourceError::Auth { .. } => SourceError::Auth { message },
332 SourceError::Refused { .. } => SourceError::Refused { message },
333 SourceError::Malformed { .. } => SourceError::Malformed { message },
334 SourceError::Unavailable { .. } => SourceError::Unavailable { message },
335 }
336}
337
338#[async_trait]
339impl TaskSource for SubprocessSource {
340 fn kind(&self) -> &'static str {
341 self.kind
342 }
343
344 fn capabilities(&self) -> Capabilities {
345 self.capabilities.clone()
346 }
347
348 async fn health(&self) -> Result<Health, SourceError> {
349 self.ask("health", json!({})).await
350 }
351
352 async fn get_task(&self, id: &NativeId) -> Result<Option<Task>, SourceError> {
353 let result: TaskResult = self
354 .ask("get_task", params(&IdParams { id: id.clone() }))
355 .await?;
356 Ok(result.task)
357 }
358
359 async fn get_project(&self, id: &NativeId) -> Result<Option<Project>, SourceError> {
360 let result: ProjectResult = self
361 .ask("get_project", params(&IdParams { id: id.clone() }))
362 .await?;
363 Ok(result.project)
364 }
365
366 async fn query_tasks(
367 &self,
368 query: &TaskQuery,
369 page: &PageRequest,
370 ) -> Result<Page<Task>, SourceError> {
371 self.ask(
372 "query_tasks",
373 params(&TaskQueryParams {
374 query: query.clone(),
375 page: page.clone(),
376 }),
377 )
378 .await
379 }
380
381 async fn query_projects(
382 &self,
383 query: &ProjectQuery,
384 page: &PageRequest,
385 ) -> Result<Page<Project>, SourceError> {
386 self.ask(
387 "query_projects",
388 params(&ProjectQueryParams {
389 query: query.clone(),
390 page: page.clone(),
391 }),
392 )
393 .await
394 }
395
396 async fn labels(&self, page: &PageRequest) -> Result<Page<Label>, SourceError> {
397 self.ask("labels", params(&LabelParams { page: page.clone() }))
398 .await
399 }
400
401 async fn task_dependencies(
402 &self,
403 id: &NativeId,
404 direction: Direction,
405 page: &PageRequest,
406 ) -> Result<Page<DependencyEdge>, SourceError> {
407 self.ask(
408 "task_dependencies",
409 params(&DependencyParams {
410 id: id.clone(),
411 direction,
412 page: page.clone(),
413 }),
414 )
415 .await
416 }
417
418 async fn project_dependencies(
419 &self,
420 id: &NativeId,
421 direction: Direction,
422 page: &PageRequest,
423 ) -> Result<Page<DependencyEdge>, SourceError> {
424 self.ask(
425 "project_dependencies",
426 params(&DependencyParams {
427 id: id.clone(),
428 direction,
429 page: page.clone(),
430 }),
431 )
432 .await
433 }
434
435 fn writes(&self) -> WriteSupport {
436 self.writes
437 }
438
439 async fn write_task(&self, write: &ItemWrite<Task>) -> Result<NativeId, SourceError> {
440 let result: WriteResult = self
441 .ask(
442 "write_task",
443 params(&TaskWriteParams {
444 write: write.clone(),
445 }),
446 )
447 .await?;
448 Ok(result.id)
449 }
450
451 async fn write_project(&self, write: &ItemWrite<Project>) -> Result<NativeId, SourceError> {
452 let result: WriteResult = self
453 .ask(
454 "write_project",
455 params(&ProjectWriteParams {
456 write: write.clone(),
457 }),
458 )
459 .await?;
460 Ok(result.id)
461 }
462
463 async fn delete_task(&self, id: &NativeId) -> Result<(), SourceError> {
464 let _: IgnoredResult = self
465 .ask("delete_task", params(&DeleteParams { id: id.clone() }))
466 .await?;
467 Ok(())
468 }
469
470 async fn delete_project(&self, id: &NativeId) -> Result<(), SourceError> {
471 let _: IgnoredResult = self
472 .ask("delete_project", params(&DeleteParams { id: id.clone() }))
473 .await?;
474 Ok(())
475 }
476
477 async fn get_document(&self, id: &NativeId) -> Result<Option<Document>, SourceError> {
478 let result: DocumentResult = self
479 .ask("get_document", params(&IdParams { id: id.clone() }))
480 .await?;
481 Ok(result.document)
482 }
483
484 async fn query_documents(
485 &self,
486 query: &DocumentQuery,
487 page: &PageRequest,
488 ) -> Result<Page<Document>, SourceError> {
489 self.ask(
490 "query_documents",
491 params(&DocumentQueryParams {
492 query: query.clone(),
493 page: page.clone(),
494 }),
495 )
496 .await
497 }
498
499 async fn write_document(&self, write: &ItemWrite<Document>) -> Result<NativeId, SourceError> {
500 let result: WriteResult = self
501 .ask(
502 "write_document",
503 params(&DocumentWriteParams {
504 write: write.clone(),
505 }),
506 )
507 .await?;
508 Ok(result.id)
509 }
510
511 async fn delete_document(&self, id: &NativeId) -> Result<(), SourceError> {
512 let _: IgnoredResult = self
513 .ask("delete_document", params(&DeleteParams { id: id.clone() }))
514 .await?;
515 Ok(())
516 }
517
518 async fn metering(&self) -> Result<Option<Metering>, SourceError> {
519 if !self.meters {
522 return Ok(None);
523 }
524 let result: MeteringResult = self.ask("metering", json!({})).await?;
525 Ok(result.metering)
526 }
527}
528
529#[derive(serde::Deserialize)]
538struct IgnoredResult {}
539
540fn params<T: serde::Serialize>(value: &T) -> Value {
545 serde_json::to_value(value).expect("method parameters are plain data")
546}