Skip to main content

stasis/application/orchestration/
handoff_pattern_pipeline.rs

1use 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}}", &current_input)
76                .replace("{input}", &current_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}