stasis/application/orchestration/
handoff_pattern_pipeline.rs1use crate::application::orchestration::prompt_pipeline::{
2 PromptExecutionContext, PromptExecutionPipeline, PromptExecutionRequest,
3};
4use crate::domain::errors::Result;
5
6#[derive(Clone, Debug)]
7pub struct HandoffPatternTurn {
8 pub actor_id: String,
9 pub user_prompt_template: String,
10 pub system_prompt: Option<String>,
11 pub policy_profile: Option<String>,
12 pub model_hint: Option<String>,
13}
14
15#[derive(Clone, Debug)]
16pub struct HandoffPatternExecutionRequest {
17 pub initial_user_prompt: String,
18 pub trace_id: Option<String>,
19 pub correlation_id: Option<String>,
20 pub policy_profile: Option<String>,
21 pub model_hint: Option<String>,
22 pub turns: Vec<HandoffPatternTurn>,
23}
24
25#[derive(Clone, Debug)]
26pub struct HandoffPatternTurnResult {
27 pub actor_id: String,
28 pub rendered_prompt: String,
29 pub output_text: String,
30}
31
32#[derive(Clone, Debug)]
33pub struct HandoffTransition {
34 pub from_actor_id: String,
35 pub to_actor_id: String,
36}
37
38#[derive(Clone, Debug)]
39pub struct HandoffPatternExecutionResponse {
40 pub final_text: String,
41 pub turns: Vec<HandoffPatternTurnResult>,
42 pub handoffs: Vec<HandoffTransition>,
43 pub termination_reason: String,
44}
45
46#[derive(Clone)]
47pub struct HandoffPatternPipeline {
48 prompt_pipeline: PromptExecutionPipeline,
49}
50
51impl HandoffPatternPipeline {
52 pub fn new(prompt_pipeline: PromptExecutionPipeline) -> Self {
53 Self { prompt_pipeline }
54 }
55
56 pub async fn execute(
57 &self,
58 request: HandoffPatternExecutionRequest,
59 ) -> Result<HandoffPatternExecutionResponse> {
60 let mut current_input = request.initial_user_prompt;
61 let mut turn_results = Vec::with_capacity(request.turns.len());
62 let mut handoffs = Vec::new();
63 let mut previous_actor: Option<String> = None;
64
65 for turn in request.turns {
66 if let Some(from_actor_id) = previous_actor.clone() {
67 handoffs.push(HandoffTransition {
68 from_actor_id,
69 to_actor_id: turn.actor_id.clone(),
70 });
71 }
72
73 let rendered_prompt = turn
74 .user_prompt_template
75 .replace("{{input}}", ¤t_input)
76 .replace("{input}", ¤t_input);
77
78 let context = PromptExecutionContext {
79 trace_id: request.trace_id.clone(),
80 correlation_id: request.correlation_id.clone(),
81 policy_profile: turn
82 .policy_profile
83 .clone()
84 .or_else(|| request.policy_profile.clone()),
85 model_hint: turn
86 .model_hint
87 .clone()
88 .or_else(|| request.model_hint.clone()),
89 };
90
91 let mut prompt_request =
92 PromptExecutionRequest::from_user_prompt(rendered_prompt.clone())
93 .with_context(context);
94 if let Some(system_prompt) = turn.system_prompt {
95 prompt_request = prompt_request.with_system_prompt(system_prompt);
96 }
97
98 let response = self.prompt_pipeline.execute(prompt_request).await?;
99 previous_actor = Some(turn.actor_id.clone());
100 current_input = response.text.clone();
101 turn_results.push(HandoffPatternTurnResult {
102 actor_id: turn.actor_id,
103 rendered_prompt,
104 output_text: response.text,
105 });
106 }
107
108 Ok(HandoffPatternExecutionResponse {
109 final_text: current_input,
110 turns: turn_results,
111 handoffs,
112 termination_reason: "completed_all_turns".to_string(),
113 })
114 }
115}