Skip to main content

stasis/application/orchestration/
sequential_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 SequentialPatternStage {
8    pub stage_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 SequentialPatternExecutionRequest {
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 stages: Vec<SequentialPatternStage>,
23}
24
25#[derive(Clone, Debug)]
26pub struct SequentialPatternStageResult {
27    pub stage_id: String,
28    pub rendered_prompt: String,
29    pub output_text: String,
30}
31
32#[derive(Clone, Debug)]
33pub struct SequentialPatternExecutionResponse {
34    pub final_text: String,
35    pub stages: Vec<SequentialPatternStageResult>,
36    pub termination_reason: String,
37}
38
39#[derive(Clone)]
40pub struct SequentialPatternPipeline {
41    prompt_pipeline: PromptExecutionPipeline,
42}
43
44impl SequentialPatternPipeline {
45    pub fn new(prompt_pipeline: PromptExecutionPipeline) -> Self {
46        Self { prompt_pipeline }
47    }
48
49    pub async fn execute(
50        &self,
51        request: SequentialPatternExecutionRequest,
52    ) -> Result<SequentialPatternExecutionResponse> {
53        let mut current_input = request.initial_user_prompt;
54        let mut stage_results = Vec::with_capacity(request.stages.len());
55
56        for stage in request.stages {
57            let rendered_prompt = stage
58                .user_prompt_template
59                .replace("{{input}}", &current_input)
60                .replace("{input}", &current_input);
61
62            let context = PromptExecutionContext {
63                trace_id: request.trace_id.clone(),
64                correlation_id: request.correlation_id.clone(),
65                policy_profile: stage
66                    .policy_profile
67                    .clone()
68                    .or_else(|| request.policy_profile.clone()),
69                model_hint: stage
70                    .model_hint
71                    .clone()
72                    .or_else(|| request.model_hint.clone()),
73            };
74
75            let mut prompt_request =
76                PromptExecutionRequest::from_user_prompt(rendered_prompt.clone())
77                    .with_context(context);
78            if let Some(system_prompt) = stage.system_prompt {
79                prompt_request = prompt_request.with_system_prompt(system_prompt);
80            }
81
82            let response = self.prompt_pipeline.execute(prompt_request).await?;
83            current_input = response.text.clone();
84            stage_results.push(SequentialPatternStageResult {
85                stage_id: stage.stage_id,
86                rendered_prompt,
87                output_text: response.text,
88            });
89        }
90
91        Ok(SequentialPatternExecutionResponse {
92            final_text: current_input,
93            stages: stage_results,
94            termination_reason: "completed_all_stages".to_string(),
95        })
96    }
97}