Skip to main content

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