Skip to main content

stasis/application/runtime/
tool_loop_job_handler.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use serde_json::{Value, json};
5
6use crate::application::orchestration::runtime_job_payloads::{
7    AgentToolCallMode, ToolLoopJobPayload,
8};
9use crate::application::runtime::chat_options_resolver::validate_reasoning_effort;
10use crate::application::runtime::identity_context_compiler::{
11    load_identity_context_summary, prepend_identity_snapshot,
12};
13use crate::application::runtime::memory_recall_context_compiler::prepend_memory_recall_context;
14use crate::application::runtime::memory_persistence_helpers::{
15    SttpPromptNodeFormat, memory_query_fingerprint, memory_query_id, memory_scope_hash,
16    render_prompt_response_sttp_node, resolve_sttp_output_node_id, should_store,
17};
18use crate::application::runtime::memory_recall_request_builder::build_memory_recall_request;
19use crate::application::orchestration::prompt_pipeline::{
20    PromptExecutionPipeline,
21};
22use crate::application::orchestration::tool_loop_pipeline::{
23    ToolCallMode, ToolLoopExecutionRequest, ToolLoopPipeline,
24};
25use crate::application::orchestration::tool_registry::ToolRegistry;
26use crate::application::runtime::in_memory_runtime::{JobExecutionOutcome, JobHandler};
27use crate::application::runtime::runtime_diagnostics_helpers::{
28    build_runtime_failure_identity_context_section, build_runtime_failure_memory_recall_section,
29    build_runtime_memory_diagnostics_bundle, RuntimeIdentityDiagnosticsInput,
30    RuntimeMemoryRecallDiagnosticsInput, RuntimeMemoryStoreDiagnosticsInput,
31};
32use crate::application::runtime::runtime_handler_execution_context::RuntimeHandlerExecutionContext;
33use crate::domain::errors::Result;
34use crate::domain::runtime::job::Job;
35use crate::ports::outbound::ai_chat_client::AiChatClient;
36use crate::ports::outbound::memory::identity_memory_store::IdentityMemoryStore;
37use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
38use crate::ports::outbound::memory::memory_context_writer::MemoryContextWriter;
39use crate::ports::outbound::memory::memory_models::MemoryStoreRequest;
40
41pub struct ToolLoopJobHandler {
42    pipeline: ToolLoopPipeline,
43    memory_reader: Option<Arc<dyn MemoryContextReader>>,
44    memory_writer: Option<Arc<dyn MemoryContextWriter>>,
45    identity_memory_store: Option<Arc<dyn IdentityMemoryStore>>,
46}
47
48impl ToolLoopJobHandler {
49    pub fn new(chat_client: Arc<dyn AiChatClient>, tool_registry: Arc<dyn ToolRegistry>) -> Self {
50        Self::new_with_memory_and_identity(chat_client, tool_registry, None, None, None)
51    }
52
53    pub fn new_with_memory(
54        chat_client: Arc<dyn AiChatClient>,
55        tool_registry: Arc<dyn ToolRegistry>,
56        memory_reader: Option<Arc<dyn MemoryContextReader>>,
57        memory_writer: Option<Arc<dyn MemoryContextWriter>>,
58    ) -> Self {
59        Self::new_with_memory_and_identity(
60            chat_client,
61            tool_registry,
62            memory_reader,
63            memory_writer,
64            None,
65        )
66    }
67
68    pub fn new_with_memory_and_identity(
69        chat_client: Arc<dyn AiChatClient>,
70        tool_registry: Arc<dyn ToolRegistry>,
71        memory_reader: Option<Arc<dyn MemoryContextReader>>,
72        memory_writer: Option<Arc<dyn MemoryContextWriter>>,
73        identity_memory_store: Option<Arc<dyn IdentityMemoryStore>>,
74    ) -> Self {
75        let prompt_pipeline = PromptExecutionPipeline::new(chat_client);
76        Self {
77            pipeline: ToolLoopPipeline::new(prompt_pipeline, tool_registry),
78            memory_reader,
79            memory_writer,
80            identity_memory_store,
81        }
82    }
83
84    fn parse_payload(raw: &str) -> std::result::Result<ToolLoopJobPayload, String> {
85        let payload: ToolLoopJobPayload = serde_json::from_str(raw)
86            .map_err(|err| format!("policy violation: invalid tool-loop payload json: {err}"))?;
87
88        if payload.user_prompt.trim().is_empty() {
89            return Err(
90                "policy violation: tool-loop payload.user_prompt must be non-empty".to_string(),
91            );
92        }
93        if payload.tool_name.trim().is_empty() {
94            return Err(
95                "policy violation: tool-loop payload.tool_name must be non-empty".to_string(),
96            );
97        }
98
99        validate_reasoning_effort(payload.reasoning_effort.as_deref())
100            .map_err(|err| format!("policy violation: {err}"))?;
101
102        Ok(payload)
103    }
104
105    fn build_failure(message: String) -> JobExecutionOutcome {
106        let diagnostics = json!({
107            "provider": "stasis-tool-loop",
108            "status": "failure",
109            "guardrail_code": "POLICY_VIOLATION",
110            "policy_reason": &message,
111        })
112        .to_string();
113
114        JobExecutionOutcome::FatalFailure {
115            message,
116            execution_id: None,
117            diagnostics: Some(diagnostics),
118        }
119    }
120
121}
122
123#[async_trait]
124impl JobHandler for ToolLoopJobHandler {
125    fn job_type(&self) -> &'static str {
126        "workflow.stasis.tool_loop"
127    }
128
129    async fn execute(&self, job: &Job) -> Result<JobExecutionOutcome> {
130        let payload = match Self::parse_payload(&job.payload_ref) {
131            Ok(payload) => payload,
132            Err(message) => return Ok(Self::build_failure(message)),
133        };
134
135        let execution_context = RuntimeHandlerExecutionContext::new(
136            job,
137            payload.policy_profile.clone(),
138            payload.model_hint.clone(),
139            payload.reasoning_effort.clone(),
140            self.memory_reader.is_some(),
141            self.memory_writer.is_some(),
142            self.identity_memory_store.is_some(),
143        );
144
145        let memory_policy = payload.memory_policy.as_ref();
146        let (identity_summary, identity_error) = load_identity_context_summary(
147            self.identity_memory_store.as_ref(),
148            execution_context.correlation_id(),
149            execution_context.policy_profile(),
150        )
151        .await;
152        let mut effective_user_prompt =
153            prepend_identity_snapshot(&payload.user_prompt, identity_summary.as_deref());
154
155        let mut memory_recall = None;
156        let mut memory_recall_error = None;
157        let mut input_memory_query_id = None;
158        let mut input_memory_query_fingerprint = None;
159        if let Some(reader) = &self.memory_reader {
160            let recall_request = build_memory_recall_request(
161                execution_context.correlation_id(),
162                Some(&effective_user_prompt),
163                memory_policy,
164            );
165            input_memory_query_id = Some(memory_query_id(
166                execution_context.correlation_id(),
167                &recall_request,
168            ));
169            input_memory_query_fingerprint = Some(memory_query_fingerprint(&recall_request));
170
171            match reader.recall(&recall_request).await {
172                Ok(response) => {
173                    effective_user_prompt = prepend_memory_recall_context(&effective_user_prompt, &response);
174                    memory_recall = Some(response);
175                }
176                Err(err) => memory_recall_error = Some(err.to_string()),
177            }
178        }
179
180        let context = execution_context.prompt_context_clone();
181
182        let request = ToolLoopExecutionRequest {
183            user_prompt: effective_user_prompt,
184            system_prompt: payload.system_prompt,
185            context,
186            tool_name: payload.tool_name,
187            tool_input: payload.tool_input.unwrap_or(Value::Null),
188            tool_call_mode: match payload.tool_call_mode {
189                Some(AgentToolCallMode::Strict) => ToolCallMode::Strict,
190                _ => ToolCallMode::Auto,
191            },
192        };
193
194        let response = match self.pipeline.execute(request).await {
195            Ok(response) => response,
196            Err(err) => {
197                let error_text = err.to_string();
198                let is_policy_violation = error_text.contains("policy violation");
199                let diagnostics = if is_policy_violation {
200                    json!({
201                        "provider": "stasis-tool-loop",
202                        "status": "failure",
203                        "guardrail_code": "POLICY_VIOLATION",
204                        "policy_reason": error_text,
205                    })
206                    .to_string()
207                } else {
208                    json!({
209                        "provider": "stasis-tool-loop",
210                        "status": "failure",
211                        "error": error_text,
212                        "memory_recall": build_runtime_failure_memory_recall_section(
213                            execution_context.memory_reader_enabled(),
214                            memory_recall_error,
215                        ),
216                        "identity_context": build_runtime_failure_identity_context_section(
217                            execution_context.identity_enabled(),
218                            identity_summary,
219                            identity_error,
220                        ),
221                    })
222                    .to_string()
223                };
224
225                return Ok(JobExecutionOutcome::FatalFailure {
226                    message: error_text,
227                    execution_id: None,
228                    diagnostics: Some(diagnostics),
229                });
230            }
231        };
232
233        let invoked_tools: Vec<String> = response
234            .tool_invocations
235            .iter()
236            .map(|invocation| invocation.tool_name.clone())
237            .collect();
238
239        let mut memory_store = None;
240        let mut memory_store_error = None;
241        if should_store(memory_policy)
242            && let Some(writer) = &self.memory_writer
243        {
244            let store_request = MemoryStoreRequest {
245                session_id: execution_context.correlation_id().to_string(),
246                raw_node: render_prompt_response_sttp_node(
247                    execution_context.correlation_id(),
248                    &response.tool_name,
249                    &response.text,
250                    SttpPromptNodeFormat::TaggedSchema,
251                ),
252            };
253
254            match writer.store_context(&store_request).await {
255                Ok(stored) => memory_store = Some(stored),
256                Err(err) => memory_store_error = Some(err.to_string()),
257            }
258        }
259
260        let sttp_output_node_id =
261            resolve_sttp_output_node_id(memory_store.as_ref(), format!("sttp:tool-loop:{}", job.id));
262        let memory_scope_hash = memory_scope_hash(execution_context.correlation_id(), memory_policy);
263        let input_memory_query_id_for_top_level = input_memory_query_id.clone();
264        let input_memory_query_fingerprint_for_top_level =
265            input_memory_query_fingerprint.clone();
266        let diagnostics_bundle = build_runtime_memory_diagnostics_bundle(
267            RuntimeMemoryRecallDiagnosticsInput {
268                attempted: execution_context.memory_reader_enabled(),
269                response: memory_recall,
270                query_id: input_memory_query_id,
271                query_fingerprint: input_memory_query_fingerprint,
272                error: memory_recall_error,
273            },
274            RuntimeMemoryStoreDiagnosticsInput {
275                attempted: execution_context.memory_writer_enabled(),
276                response: memory_store,
277                error: memory_store_error,
278            },
279            RuntimeIdentityDiagnosticsInput {
280                attempted: execution_context.identity_enabled(),
281                summary: identity_summary,
282                error: identity_error,
283            },
284        );
285
286        let diagnostics = json!({
287            "provider": "stasis-tool-loop",
288            "status": "success",
289            "tool_name": response.tool_name,
290            "tool_output": response.tool_output,
291            "tool_invocations": response.tool_invocations,
292            "invoked_tools": invoked_tools,
293            "tool_rounds": response.rounds_executed,
294            "termination_reason": response.termination_reason,
295            "policy_profile": response.metadata.policy_profile,
296            "model_hint": response.metadata.model_hint,
297            "output_preview": response.text.chars().take(160).collect::<String>(),
298            "memory_retrieved_count": diagnostics_bundle.retrieved_count,
299            "memory_retrieval_path": diagnostics_bundle.retrieval_path,
300            "memory_fallback_triggered": diagnostics_bundle.fallback_triggered,
301            "memory_fallback_reason": diagnostics_bundle.fallback_reason,
302            "memory_scope_hash": memory_scope_hash,
303            "memory_store_valid": diagnostics_bundle.store_valid,
304            "memory_store_node_id": diagnostics_bundle.store_node_id,
305            "input_memory_query_id": input_memory_query_id_for_top_level,
306            "input_memory_query_fingerprint": input_memory_query_fingerprint_for_top_level,
307            "output_memory_node_id": diagnostics_bundle.store_node_id,
308            "memory_recall": diagnostics_bundle.memory_recall,
309            "identity_context": diagnostics_bundle.identity_context,
310            "memory_store": diagnostics_bundle.memory_store,
311        })
312        .to_string();
313
314        Ok(JobExecutionOutcome::Success {
315            sttp_output_node_id,
316            execution_id: None,
317            diagnostics: Some(diagnostics),
318        })
319    }
320}