Skip to main content

stasis/application/orchestration/
concurrent_pattern_pipeline.rs

1use std::sync::Arc;
2
3use serde_json::Value;
4
5use crate::application::orchestration::prompt_pipeline::{
6    PromptExecutionContext, PromptExecutionPipeline, PromptExecutionRequest,
7};
8use crate::application::orchestration::runtime_job_payloads::{
9    ConcurrentBranchExecutionMode, MemoryPolicyPayload,
10};
11use crate::application::orchestration::tool_loop_pipeline::{
12    ToolCallMode, ToolInvocation, ToolLoopExecutionRequest, ToolLoopPipeline,
13};
14use crate::application::orchestration::tool_registry::ToolRegistry;
15use crate::application::runtime::concurrent_tool_branch_memory::{
16    prepare_concurrent_tool_branch, store_concurrent_tool_branch_memory,
17};
18use crate::application::runtime::chat_options_resolver::resolve_reasoning_effort;
19use crate::domain::errors::{Result, StasisError};
20use crate::ports::outbound::memory::identity_memory_store::IdentityMemoryStore;
21use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
22use crate::ports::outbound::memory::memory_context_writer::MemoryContextWriter;
23use tokio::task::JoinSet;
24
25#[derive(Clone, Debug)]
26pub struct ConcurrentPatternBranch {
27    pub branch_id: String,
28    pub user_prompt_template: String,
29    pub system_prompt: Option<String>,
30    pub policy_profile: Option<String>,
31    pub model_hint: Option<String>,
32    pub reasoning_effort: Option<String>,
33    pub execution_mode: ConcurrentBranchExecutionMode,
34    pub tool_name: Option<String>,
35    pub tool_input: Option<Value>,
36    pub tool_call_mode: ToolCallMode,
37    pub memory_policy: Option<MemoryPolicyPayload>,
38}
39
40#[derive(Clone, Debug)]
41pub struct ConcurrentPatternExecutionRequest {
42    pub initial_user_prompt: String,
43    pub trace_id: Option<String>,
44    pub correlation_id: Option<String>,
45    pub policy_profile: Option<String>,
46    pub model_hint: Option<String>,
47    pub reasoning_effort: Option<String>,
48    pub default_memory_policy: Option<MemoryPolicyPayload>,
49    pub merge_strategy: Option<String>,
50    pub branches: Vec<ConcurrentPatternBranch>,
51}
52
53#[derive(Clone, Debug)]
54pub struct ConcurrentPatternBranchResult {
55    pub branch_id: String,
56    pub execution_mode: ConcurrentBranchExecutionMode,
57    pub rendered_prompt: String,
58    pub output_text: String,
59    pub tool_name: Option<String>,
60    pub tool_output: Option<Value>,
61    pub tool_invocations: Vec<ToolInvocation>,
62    pub rounds_executed: Option<usize>,
63    pub branch_termination_reason: Option<String>,
64    pub memory_retrieved_count: Option<usize>,
65    pub memory_store_node_id: Option<String>,
66    pub input_memory_query_id: Option<String>,
67    pub input_memory_query_fingerprint: Option<String>,
68    pub memory_recall_error: Option<String>,
69    pub memory_store_error: Option<String>,
70}
71
72#[derive(Clone, Debug)]
73pub struct ConcurrentPatternExecutionResponse {
74    pub final_text: String,
75    pub branches: Vec<ConcurrentPatternBranchResult>,
76    pub termination_reason: String,
77    pub merge_strategy: String,
78}
79
80#[derive(Clone)]
81pub struct ConcurrentPatternPipeline {
82    prompt_pipeline: PromptExecutionPipeline,
83    tool_loop_pipeline: Option<ToolLoopPipeline>,
84    memory_reader: Option<Arc<dyn MemoryContextReader>>,
85    memory_writer: Option<Arc<dyn MemoryContextWriter>>,
86    identity_memory_store: Option<Arc<dyn IdentityMemoryStore>>,
87}
88
89#[derive(Clone)]
90struct ConcurrentSharedInputs {
91    initial_input: Arc<str>,
92    trace_id: Arc<Option<String>>,
93    correlation_id: Arc<Option<String>>,
94    default_policy_profile: Arc<Option<String>>,
95    default_model_hint: Arc<Option<String>>,
96    default_reasoning_effort: Arc<Option<String>>,
97    default_memory_policy: Arc<Option<MemoryPolicyPayload>>,
98}
99
100impl ConcurrentSharedInputs {
101    fn build_context(
102        &self,
103        policy_profile: Option<String>,
104        model_hint: Option<String>,
105        reasoning_effort: Option<String>,
106    ) -> PromptExecutionContext {
107        PromptExecutionContext {
108            trace_id: (*self.trace_id).clone(),
109            correlation_id: (*self.correlation_id).clone(),
110            policy_profile: policy_profile.or_else(|| (*self.default_policy_profile).clone()),
111            model_hint: model_hint.or_else(|| (*self.default_model_hint).clone()),
112            reasoning_effort: resolve_reasoning_effort(
113                reasoning_effort,
114                (*self.default_reasoning_effort).clone(),
115            ),
116        }
117    }
118
119    fn render_template(&self, template: &str) -> String {
120        template
121            .replace("{{input}}", &self.initial_input)
122            .replace("{input}", &self.initial_input)
123    }
124}
125
126impl ConcurrentPatternPipeline {
127    pub fn new(prompt_pipeline: PromptExecutionPipeline) -> Self {
128        Self {
129            prompt_pipeline,
130            tool_loop_pipeline: None,
131            memory_reader: None,
132            memory_writer: None,
133            identity_memory_store: None,
134        }
135    }
136
137    pub fn new_with_tool_loop(
138        prompt_pipeline: PromptExecutionPipeline,
139        tool_registry: Arc<dyn ToolRegistry>,
140        memory_reader: Option<Arc<dyn MemoryContextReader>>,
141        memory_writer: Option<Arc<dyn MemoryContextWriter>>,
142        identity_memory_store: Option<Arc<dyn IdentityMemoryStore>>,
143    ) -> Self {
144        Self {
145            tool_loop_pipeline: Some(ToolLoopPipeline::new(
146                prompt_pipeline.clone(),
147                tool_registry,
148            )),
149            prompt_pipeline,
150            memory_reader,
151            memory_writer,
152            identity_memory_store,
153        }
154    }
155
156    pub async fn execute(
157        &self,
158        request: ConcurrentPatternExecutionRequest,
159    ) -> Result<ConcurrentPatternExecutionResponse> {
160        let ConcurrentPatternExecutionRequest {
161            initial_user_prompt,
162            trace_id,
163            correlation_id,
164            policy_profile,
165            model_hint,
166            reasoning_effort,
167            default_memory_policy,
168            merge_strategy,
169            branches,
170        } = request;
171
172        let merge_strategy = merge_strategy.unwrap_or_else(|| "join_with_headers".to_string());
173        let shared_inputs = ConcurrentSharedInputs {
174            initial_input: Arc::<str>::from(initial_user_prompt),
175            trace_id: Arc::new(trace_id),
176            correlation_id: Arc::new(correlation_id),
177            default_policy_profile: Arc::new(policy_profile),
178            default_model_hint: Arc::new(model_hint),
179            default_reasoning_effort: Arc::new(reasoning_effort),
180            default_memory_policy: Arc::new(default_memory_policy),
181        };
182
183        let mut join_set: JoinSet<Result<ConcurrentPatternBranchResult>> = JoinSet::new();
184
185        for branch in branches {
186            let prompt_pipeline = self.prompt_pipeline.clone();
187            let tool_loop_pipeline = self.tool_loop_pipeline.clone();
188            let memory_reader = self.memory_reader.clone();
189            let memory_writer = self.memory_writer.clone();
190            let identity_memory_store = self.identity_memory_store.clone();
191            let shared_inputs = shared_inputs.clone();
192
193            join_set.spawn(async move {
194                let ConcurrentPatternBranch {
195                    branch_id,
196                    user_prompt_template,
197                    system_prompt,
198                    policy_profile,
199                    model_hint,
200                    reasoning_effort,
201                    execution_mode,
202                    tool_name,
203                    tool_input,
204                    tool_call_mode,
205                    memory_policy,
206                } = branch;
207
208                let rendered_prompt = shared_inputs.render_template(&user_prompt_template);
209                let context = shared_inputs.build_context(
210                    policy_profile.clone(),
211                    model_hint,
212                    reasoning_effort,
213                );
214                let correlation_id = shared_inputs
215                    .correlation_id
216                    .as_deref()
217                    .map(str::to_string)
218                    .unwrap_or_else(|| "unknown".to_string());
219                let resolved_memory_policy = memory_policy
220                    .or_else(|| (*shared_inputs.default_memory_policy).clone());
221                let memory_policy_ref = resolved_memory_policy.as_ref();
222
223                match execution_mode {
224                    ConcurrentBranchExecutionMode::Prompt => {
225                        let mut prompt_request =
226                            PromptExecutionRequest::from_user_prompt(rendered_prompt.clone())
227                                .with_context(context);
228                        if let Some(system_prompt) = system_prompt {
229                            prompt_request = prompt_request.with_system_prompt(system_prompt);
230                        }
231
232                        let response = prompt_pipeline.execute(prompt_request).await?;
233                        Ok(ConcurrentPatternBranchResult {
234                            branch_id,
235                            execution_mode,
236                            rendered_prompt,
237                            output_text: response.text,
238                            tool_name: None,
239                            tool_output: None,
240                            tool_invocations: Vec::new(),
241                            rounds_executed: None,
242                            branch_termination_reason: None,
243                            memory_retrieved_count: None,
244                            memory_store_node_id: None,
245                            input_memory_query_id: None,
246                            input_memory_query_fingerprint: None,
247                            memory_recall_error: None,
248                            memory_store_error: None,
249                        })
250                    }
251                    ConcurrentBranchExecutionMode::ToolLoop => {
252                        let Some(tool_loop_pipeline) = tool_loop_pipeline else {
253                            return Err(StasisError::PortFailure(
254                                "concurrent pattern tool_loop branch requires a tool registry"
255                                    .to_string(),
256                            ));
257                        };
258
259                        let tool_name = tool_name.unwrap_or_default();
260                        let tool_input =
261                            tool_input.unwrap_or_else(|| Value::Object(Default::default()));
262
263                        let prepared = prepare_concurrent_tool_branch(
264                            memory_reader.as_ref(),
265                            identity_memory_store.as_ref(),
266                            &correlation_id,
267                            context.policy_profile.as_deref(),
268                            &rendered_prompt,
269                            memory_policy_ref,
270                        )
271                        .await;
272
273                        let tool_loop_request = ToolLoopExecutionRequest {
274                            user_prompt: prepared.user_prompt,
275                            system_prompt,
276                            context,
277                            tool_name: tool_name.clone(),
278                            tool_input,
279                            tool_call_mode,
280                        };
281
282                        let response = tool_loop_pipeline.execute(tool_loop_request).await?;
283
284                        let stored = store_concurrent_tool_branch_memory(
285                            memory_writer.as_ref(),
286                            &correlation_id,
287                            &branch_id,
288                            &response.tool_name,
289                            &response.text,
290                            memory_policy_ref,
291                        )
292                        .await;
293
294                        Ok(ConcurrentPatternBranchResult {
295                            branch_id,
296                            execution_mode,
297                            rendered_prompt,
298                            output_text: response.text,
299                            tool_name: Some(response.tool_name),
300                            tool_output: Some(response.tool_output),
301                            tool_invocations: response.tool_invocations,
302                            rounds_executed: Some(response.rounds_executed),
303                            branch_termination_reason: Some(response.termination_reason),
304                            memory_retrieved_count: prepared
305                                .memory_recall
306                                .as_ref()
307                                .map(|recall| recall.retrieved),
308                            memory_store_node_id: stored
309                                .memory_store
310                                .as_ref()
311                                .map(|store| store.node_id.clone()),
312                            input_memory_query_id: prepared.input_memory_query_id,
313                            input_memory_query_fingerprint: prepared.input_memory_query_fingerprint,
314                            memory_recall_error: prepared.memory_recall_error,
315                            memory_store_error: stored.memory_store_error,
316                        })
317                    }
318                }
319            });
320        }
321
322        let mut results = Vec::new();
323        while let Some(joined) = join_set.join_next().await {
324            let result = joined.map_err(|err| {
325                StasisError::PortFailure(format!("concurrent pattern join failure: {err}"))
326            })??;
327            results.push(result);
328        }
329
330        results.sort_by(|a, b| a.branch_id.cmp(&b.branch_id));
331
332        let final_text = render_final_text(&results, &merge_strategy);
333
334        Ok(ConcurrentPatternExecutionResponse {
335            final_text,
336            branches: results,
337            termination_reason: "completed_all_branches".to_string(),
338            merge_strategy,
339        })
340    }
341}
342
343fn render_final_text(results: &[ConcurrentPatternBranchResult], merge_strategy: &str) -> String {
344    if results.is_empty() {
345        return String::new();
346    }
347
348    match merge_strategy {
349        "join_lines" => {
350            let total_text_len: usize = results.iter().map(|branch| branch.output_text.len()).sum();
351            let mut final_text = String::with_capacity(total_text_len + results.len().saturating_sub(1));
352            for (idx, branch) in results.iter().enumerate() {
353                if idx > 0 {
354                    final_text.push('\n');
355                }
356                final_text.push_str(&branch.output_text);
357            }
358            final_text
359        }
360        _ => {
361            let total_text_len: usize = results
362                .iter()
363                .map(|branch| branch.branch_id.len() + branch.output_text.len() + 4)
364                .sum();
365            let separator_len = 2 * results.len().saturating_sub(1);
366            let mut final_text = String::with_capacity(total_text_len + separator_len);
367            for (idx, branch) in results.iter().enumerate() {
368                if idx > 0 {
369                    final_text.push_str("\n\n");
370                }
371                final_text.push('[');
372                final_text.push_str(&branch.branch_id);
373                final_text.push_str("]\n");
374                final_text.push_str(&branch.output_text);
375            }
376            final_text
377        }
378    }
379}
380
381#[cfg(test)]
382mod tests {
383    use std::sync::Arc;
384
385    use async_trait::async_trait;
386    use genai::adapter::AdapterKind;
387    use genai::ModelIden;
388    use genai::chat::{
389        ChatOptions, ChatRequest, ChatResponse, MessageContent, ToolCall, Usage,
390    };
391    use serde_json::json;
392
393    use super::*;
394    use crate::application::orchestration::tool_registry::{InMemoryToolRegistry, StasisTool};
395    use crate::domain::errors::Result as StasisResult;
396    use crate::ports::outbound::ai_chat_client::AiChatClient;
397
398    struct EchoPromptChatClient;
399
400    #[async_trait]
401    impl AiChatClient for EchoPromptChatClient {
402        async fn complete(
403            &self,
404            request: ChatRequest,
405            _options: Option<&ChatOptions>,
406        ) -> StasisResult<ChatResponse> {
407            let echoed_text = request
408                .messages
409                .iter()
410                .rev()
411                .filter_map(|message| message.content.first_text())
412                .next()
413                .unwrap_or_default();
414
415            Ok(ChatResponse {
416                content: MessageContent::from_text(format!("echo::{echoed_text}")),
417                reasoning_content: None,
418                model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
419                provider_model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
420                stop_reason: None,
421                usage: Usage::default(),
422                captured_raw_body: None,
423                response_id: None,
424            })
425        }
426    }
427
428    struct BranchAwareToolCallClient;
429
430    #[async_trait]
431    impl AiChatClient for BranchAwareToolCallClient {
432        async fn complete(
433            &self,
434            request: ChatRequest,
435            _options: Option<&ChatOptions>,
436        ) -> StasisResult<ChatResponse> {
437            let user_text = request
438                .messages
439                .iter()
440                .rev()
441                .filter_map(|message| message.content.first_text())
442                .next()
443                .unwrap_or_default();
444
445            if user_text.contains("Tool branch") {
446                let has_tool_response = request
447                    .messages
448                    .iter()
449                    .any(|message| !message.content.tool_responses().is_empty());
450
451                if !has_tool_response {
452                    return Ok(ChatResponse {
453                        content: MessageContent::from_tool_calls(vec![ToolCall {
454                            call_id: "tool-call-1".to_string(),
455                            fn_name: "stasis.web.search.mock".to_string(),
456                            fn_arguments: json!({ "query": "branch query" }),
457                            thought_signatures: None,
458                        }]),
459                        reasoning_content: None,
460                        model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
461                        provider_model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
462                        stop_reason: None,
463                        usage: Usage::default(),
464                        captured_raw_body: None,
465                        response_id: None,
466                    });
467                }
468
469                return Ok(ChatResponse {
470                    content: MessageContent::from_text("tool branch final answer"),
471                    reasoning_content: None,
472                    model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
473                    provider_model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
474                    stop_reason: None,
475                    usage: Usage::default(),
476                    captured_raw_body: None,
477                    response_id: None,
478                });
479            }
480
481            Ok(ChatResponse {
482                content: MessageContent::from_text(format!("echo::{user_text}")),
483                reasoning_content: None,
484                model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
485                provider_model_iden: ModelIden::new(AdapterKind::OpenAI, "gpt-4o-mini"),
486                stop_reason: None,
487                usage: Usage::default(),
488                captured_raw_body: None,
489                response_id: None,
490            })
491        }
492    }
493
494    struct MockWebSearchTool;
495
496    #[async_trait]
497    impl StasisTool for MockWebSearchTool {
498        fn name(&self) -> &'static str {
499            "stasis.web.search.mock"
500        }
501
502        async fn invoke(&self, input: Value) -> StasisResult<Value> {
503            Ok(json!({
504                "query": input.get("query").cloned().unwrap_or(json!("unknown")),
505                "results": [{"title": "mock result"}]
506            }))
507        }
508    }
509
510    #[tokio::test]
511    async fn concurrent_pattern_mixed_branches_execute() {
512        let chat_client = Arc::new(BranchAwareToolCallClient);
513        let prompt_pipeline = PromptExecutionPipeline::new(chat_client);
514        let tool_registry = Arc::new(InMemoryToolRegistry::default());
515        tool_registry
516            .register_tool(MockWebSearchTool)
517            .expect("tool should register");
518
519        let pipeline = ConcurrentPatternPipeline::new_with_tool_loop(
520            prompt_pipeline,
521            tool_registry,
522            None,
523            None,
524            None,
525        );
526
527        let response = pipeline
528            .execute(ConcurrentPatternExecutionRequest {
529                initial_user_prompt: "shared input".to_string(),
530                trace_id: None,
531                correlation_id: None,
532                policy_profile: None,
533                model_hint: None,
534                reasoning_effort: None,
535                default_memory_policy: None,
536                merge_strategy: Some("join_with_headers".to_string()),
537                branches: vec![
538                    ConcurrentPatternBranch {
539                        branch_id: "prompt".to_string(),
540                        user_prompt_template: "Prompt branch {{input}}".to_string(),
541                        system_prompt: None,
542                        policy_profile: None,
543                        model_hint: None,
544                        reasoning_effort: None,
545                        execution_mode: ConcurrentBranchExecutionMode::Prompt,
546                        tool_name: None,
547                        tool_input: None,
548                        tool_call_mode: ToolCallMode::Auto,
549                        memory_policy: None,
550                    },
551                    ConcurrentPatternBranch {
552                        branch_id: "tool".to_string(),
553                        user_prompt_template: "Tool branch {{input}}".to_string(),
554                        system_prompt: None,
555                        policy_profile: None,
556                        model_hint: None,
557                        reasoning_effort: None,
558                        execution_mode: ConcurrentBranchExecutionMode::ToolLoop,
559                        tool_name: Some("stasis.web.search.mock".to_string()),
560                        tool_input: Some(json!({ "query": "shared input" })),
561                        tool_call_mode: ToolCallMode::Auto,
562                        memory_policy: None,
563                    },
564                ],
565            })
566            .await
567            .expect("mixed concurrent pattern should succeed");
568
569        assert_eq!(response.branches.len(), 2);
570        assert_eq!(response.branches[0].branch_id, "prompt");
571        assert_eq!(
572            response.branches[0].output_text,
573            "echo::Prompt branch shared input"
574        );
575        assert_eq!(response.branches[1].branch_id, "tool");
576        assert_eq!(
577            response.branches[1].output_text,
578            "tool branch final answer"
579        );
580        assert_eq!(
581            response.branches[1].execution_mode,
582            ConcurrentBranchExecutionMode::ToolLoop
583        );
584        assert_eq!(response.branches[1].tool_invocations.len(), 1);
585        assert!(response.final_text.contains("[prompt]"));
586        assert!(response.final_text.contains("[tool]"));
587    }
588
589    #[tokio::test]
590    async fn concurrent_pattern_prompt_only_without_tool_registry() {
591        let chat_client = Arc::new(EchoPromptChatClient);
592        let pipeline = ConcurrentPatternPipeline::new(PromptExecutionPipeline::new(chat_client));
593
594        let response = pipeline
595            .execute(ConcurrentPatternExecutionRequest {
596                initial_user_prompt: "base".to_string(),
597                trace_id: None,
598                correlation_id: None,
599                policy_profile: None,
600                model_hint: None,
601                reasoning_effort: None,
602                default_memory_policy: None,
603                merge_strategy: None,
604                branches: vec![ConcurrentPatternBranch {
605                    branch_id: "alpha".to_string(),
606                    user_prompt_template: "Branch {input}".to_string(),
607                    system_prompt: None,
608                    policy_profile: None,
609                    model_hint: None,
610                    reasoning_effort: None,
611                    execution_mode: ConcurrentBranchExecutionMode::Prompt,
612                    tool_name: None,
613                    tool_input: None,
614                    tool_call_mode: ToolCallMode::Auto,
615                    memory_policy: None,
616                }],
617            })
618            .await
619            .expect("prompt-only concurrent pattern should succeed");
620
621        assert_eq!(response.branches[0].output_text, "echo::Branch base");
622    }
623}