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