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}