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