Skip to main content

stasis/application/runtime/
concurrent_tool_branch_memory.rs

1use std::sync::Arc;
2
3use crate::application::orchestration::runtime_job_payloads::MemoryPolicyPayload;
4use crate::application::runtime::identity_context_compiler::{
5    load_identity_context_summary, prepend_identity_snapshot,
6};
7use crate::application::runtime::memory_persistence_helpers::{
8    SttpPromptNodeFormat, memory_query_fingerprint, memory_query_id, should_store,
9    render_prompt_response_sttp_node,
10};
11use crate::application::runtime::memory_recall_context_compiler::prepend_memory_recall_context;
12use crate::application::runtime::memory_recall_request_builder::build_memory_recall_request;
13use crate::ports::outbound::memory::identity_memory_store::IdentityMemoryStore;
14use crate::ports::outbound::memory::memory_context_reader::MemoryContextReader;
15use crate::ports::outbound::memory::memory_context_writer::MemoryContextWriter;
16use crate::ports::outbound::memory::memory_models::{
17    MemoryRecallResponse, MemoryStoreRequest, MemoryStoreResponse,
18};
19
20#[derive(Clone, Debug, Default)]
21pub struct PreparedConcurrentToolBranch {
22    pub user_prompt: String,
23    pub memory_recall: Option<MemoryRecallResponse>,
24    pub memory_recall_error: Option<String>,
25    pub identity_summary: Option<String>,
26    pub identity_error: Option<String>,
27    pub input_memory_query_id: Option<String>,
28    pub input_memory_query_fingerprint: Option<String>,
29}
30
31#[derive(Clone, Debug, Default)]
32pub struct StoredConcurrentToolBranchMemory {
33    pub memory_store: Option<MemoryStoreResponse>,
34    pub memory_store_error: Option<String>,
35}
36
37pub fn branch_memory_session_id(correlation_id: &str, branch_id: &str) -> String {
38    format!("{correlation_id}::concurrent-branch::{branch_id}")
39}
40
41pub async fn prepare_concurrent_tool_branch(
42    memory_reader: Option<&Arc<dyn MemoryContextReader>>,
43    identity_memory_store: Option<&Arc<dyn IdentityMemoryStore>>,
44    correlation_id: &str,
45    policy_profile: Option<&str>,
46    rendered_prompt: &str,
47    memory_policy: Option<&MemoryPolicyPayload>,
48) -> PreparedConcurrentToolBranch {
49    let (identity_summary, identity_error) = load_identity_context_summary(
50        identity_memory_store,
51        correlation_id,
52        policy_profile,
53    )
54    .await;
55
56    let mut effective_user_prompt =
57        prepend_identity_snapshot(rendered_prompt, identity_summary.as_deref());
58
59    let mut memory_recall = None;
60    let mut memory_recall_error = None;
61    let mut input_memory_query_id = None;
62    let mut input_memory_query_fingerprint = None;
63
64    if let Some(reader) = memory_reader {
65        let recall_request = build_memory_recall_request(
66            correlation_id,
67            Some(&effective_user_prompt),
68            memory_policy,
69        );
70        input_memory_query_id = Some(memory_query_id(correlation_id, &recall_request));
71        input_memory_query_fingerprint = Some(memory_query_fingerprint(&recall_request));
72
73        match reader.recall(&recall_request).await {
74            Ok(response) => {
75                effective_user_prompt = prepend_memory_recall_context(&effective_user_prompt, &response);
76                memory_recall = Some(response);
77            }
78            Err(err) => memory_recall_error = Some(err.to_string()),
79        }
80    }
81
82    PreparedConcurrentToolBranch {
83        user_prompt: effective_user_prompt,
84        memory_recall,
85        memory_recall_error,
86        identity_summary,
87        identity_error,
88        input_memory_query_id,
89        input_memory_query_fingerprint,
90    }
91}
92
93pub async fn store_concurrent_tool_branch_memory(
94    memory_writer: Option<&Arc<dyn MemoryContextWriter>>,
95    correlation_id: &str,
96    branch_id: &str,
97    tool_name: &str,
98    response_text: &str,
99    memory_policy: Option<&MemoryPolicyPayload>,
100) -> StoredConcurrentToolBranchMemory {
101    if !should_store(memory_policy) {
102        return StoredConcurrentToolBranchMemory::default();
103    }
104
105    let Some(writer) = memory_writer else {
106        return StoredConcurrentToolBranchMemory::default();
107    };
108
109    let session_id = branch_memory_session_id(correlation_id, branch_id);
110    let store_request = MemoryStoreRequest {
111        session_id,
112        raw_node: render_prompt_response_sttp_node(
113            correlation_id,
114            tool_name,
115            response_text,
116            SttpPromptNodeFormat::TaggedSchema,
117        ),
118    };
119
120    match writer.store_context(&store_request).await {
121        Ok(stored) => StoredConcurrentToolBranchMemory {
122            memory_store: Some(stored),
123            memory_store_error: None,
124        },
125        Err(err) => StoredConcurrentToolBranchMemory {
126            memory_store: None,
127            memory_store_error: Some(err.to_string()),
128        },
129    }
130}