Skip to main content

atman_runtime/tools/
session.rs

1use crate::error::RuntimeError;
2use crate::message::{Message, MessageRole};
3use crate::tool::{BoxFut, Tier, Tool, ToolArgs, ToolCtx, ToolResult};
4use crate::value::Value;
5
6pub struct SessionPush;
7
8impl Tool for SessionPush {
9    fn name(&self) -> &str {
10        "session.push"
11    }
12
13    fn tier(&self) -> Tier {
14        Tier::Zero
15    }
16
17    fn description(&self) -> Option<&str> {
18        Some(
19            "Push a Message value into the current session's message history. \
20             Use after dispatch_all to persist tool results so the next \
21             llm.call(context: \"session\") call can see them. The message role \
22             (user/assistant/tool/system) is preserved. Returns unit.",
23        )
24    }
25
26    fn input_schema(&self) -> serde_json::Value {
27        serde_json::json!({
28            "type": "object",
29            "properties": {
30                "message": {
31                    "type": "object",
32                    "description": "The Message value to push (e.g. a tool_result from dispatch_all). Pass the value returned by dispatch_all directly."
33                }
34            },
35            "required": ["message"]
36        })
37    }
38
39    fn call<'a>(&'a self, args: ToolArgs, ctx: &'a ToolCtx) -> BoxFut<'a, ToolResult> {
40        Box::pin(async move {
41            let val = match args.named("message").or_else(|| args.positional(0).ok()) {
42                Some(v) => v.clone(),
43                None => {
44                    return Err(RuntimeError::MissingArg("session.push: message".into()));
45                }
46            };
47            let msgs = match val {
48                Value::Message(m) => vec![m],
49                Value::List(items) => items
50                    .into_iter()
51                    .filter_map(|v| match v {
52                        Value::Message(m) => Some(m),
53                        _ => None,
54                    })
55                    .collect(),
56                other => {
57                    return Err(RuntimeError::TypeMismatch {
58                        expected: "message or list of message".into(),
59                        actual: other.kind_name().into(),
60                    });
61                }
62            };
63            let Some(_handle) = &ctx.session_messages_handle else {
64                return Err(RuntimeError::ToolFailed(
65                    "session.push: no session messages handle available".into(),
66                ));
67            };
68            let _compact_guard = match &ctx.compact_lock_handle {
69                Some(lock) => Some(lock.lock().await),
70                None => None,
71            };
72            for msg in msgs {
73                let msg = crate::tools::tool_output::maybe_truncate_tool_message_with_budget(
74                    &msg,
75                    ctx.output_store.as_deref(),
76                    ctx.tool_output_budget,
77                );
78                append_message_to_context(ctx, msg)?;
79            }
80            Ok(Value::Unit)
81        })
82    }
83}
84
85pub(crate) fn append_message_to_context(ctx: &ToolCtx, msg: Message) -> Result<(), RuntimeError> {
86    let Some(handle) = &ctx.session_messages_handle else {
87        return Err(RuntimeError::ToolFailed(
88            "session message context is unavailable".into(),
89        ));
90    };
91    emit_message_event(ctx, &msg);
92    let flow_run_id = match msg.role {
93        MessageRole::Assistant | MessageRole::Tool => {
94            ctx.flow_run_id.as_ref().map(|run_id| run_id.0.to_string())
95        }
96        MessageRole::User => ctx.message_flow_run_id().map(|run_id| run_id.0.to_string()),
97        MessageRole::System => None,
98    };
99    if msg.origin != crate::message::MessageOrigin::Internal
100        && msg.role != MessageRole::System
101        && let Some(tx) = &ctx.stream_tx
102    {
103        let _ = tx.send(crate::stream::StreamFrame::ToolResultMsg {
104            flow_run_id,
105            message: msg.clone(),
106        });
107    }
108    let mut messages = handle.lock().unwrap();
109    messages.push(msg);
110    Ok(())
111}
112
113fn emit_message_event(ctx: &ToolCtx, msg: &Message) {
114    use crate::event::{Event, TurnId};
115    let Some(sink) = &ctx.events else {
116        return;
117    };
118    let turn_id = ctx.turn_id.clone().unwrap_or_else(TurnId::now);
119    let flow_run_id = ctx.message_flow_run_id();
120    let event = match msg.role {
121        MessageRole::User => Event::UserMsg {
122            turn_id,
123            flow_run_id,
124            message: msg.clone(),
125        },
126        MessageRole::Assistant => Event::AssistantMsg {
127            turn_id,
128            flow_run_id,
129            message: msg.clone(),
130        },
131        MessageRole::Tool => Event::ToolResultMsg {
132            turn_id,
133            flow_run_id,
134            message: msg.clone(),
135        },
136        MessageRole::System => Event::SystemMsg {
137            turn_id,
138            flow_run_id: ctx.message_flow_run_id(),
139            message: msg.clone(),
140        },
141    };
142    sink.emit(event);
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148
149    #[test]
150    fn spawned_context_does_not_rewrite_messages_outside_compaction() {
151        let messages = std::sync::Arc::new(std::sync::Mutex::new(
152            (0..100)
153                .map(|index| {
154                    Message::assistant_text(crate::event::TurnId::now(), format!("old-{index}"))
155                })
156                .collect(),
157        ));
158        let mut ctx = ToolCtx::new()
159            .with_history_segment(crate::tool::HistorySegment::Spawned)
160            .with_session_messages_handle(std::sync::Arc::clone(&messages));
161        ctx.context_epoch_handle = Some(std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0)));
162
163        append_message_to_context(
164            &ctx,
165            Message::assistant_text(crate::event::TurnId::now(), "new"),
166        )
167        .unwrap();
168
169        assert_eq!(messages.lock().unwrap().len(), 101);
170        assert_eq!(messages.lock().unwrap()[0].text_concat(), "old-0");
171        assert_eq!(ctx.context_epoch_seed().as_deref(), Some("generation:0"));
172    }
173
174    #[test]
175    fn session_push_name_and_tier() {
176        let tool = SessionPush;
177        assert_eq!(tool.name(), "session.push");
178        assert_eq!(tool.tier(), Tier::Zero);
179        assert!(tool.description().is_some());
180    }
181
182    #[tokio::test]
183    async fn session_push_scopes_live_tool_result_without_scoping_session_history() {
184        use crate::event::{Event, FlowRunId, TurnId};
185        use crate::message::{MessageOrigin, MessagePart};
186
187        let session = std::sync::Arc::new(crate::session::Session::open_ephemeral());
188        let run_id = FlowRunId::now();
189        let (stream_tx, mut stream_rx) = tokio::sync::broadcast::channel(8);
190        let ctx = ToolCtx::new()
191            .with_anchors(Some(TurnId::now()), Some(run_id.clone()), None)
192            .with_events(session.sink().clone())
193            .with_session_messages_handle(session.messages_handle())
194            .with_session_runtime(session.clone())
195            .with_stream_tx(stream_tx);
196        let message = Message {
197            role: MessageRole::Tool,
198            parts: vec![MessagePart::ToolResult {
199                tool_use_id: "tu_1".into(),
200                content: "done".into(),
201                is_error: false,
202            }],
203            turn_id: TurnId::now(),
204            origin: MessageOrigin::User,
205        };
206
207        SessionPush
208            .call(
209                ToolArgs {
210                    positional: vec![Value::Message(message)],
211                    named: Vec::new(),
212                },
213                &ctx,
214            )
215            .await
216            .unwrap();
217
218        let crate::stream::StreamFrame::ToolResultMsg { flow_run_id, .. } =
219            stream_rx.recv().await.unwrap()
220        else {
221            panic!("tool result frame");
222        };
223        assert_eq!(flow_run_id.as_deref(), Some(run_id.0.to_string().as_str()));
224        assert!(session.sink().snapshot().iter().any(|event| {
225            matches!(
226                event,
227                Event::ToolResultMsg {
228                    flow_run_id: None,
229                    ..
230                }
231            )
232        }));
233        assert!(
234            session
235                .messages()
236                .iter()
237                .any(|message| message.role == MessageRole::Tool)
238        );
239    }
240
241    #[test]
242    fn internal_child_system_message_is_audited_without_becoming_a_live_transcript_frame() {
243        use crate::context_plan::{
244            ContextRecord, ContextRecordAuthority, ContextRecordBody, ContextRecordRetention,
245        };
246        use crate::event::{Event, FlowRunId, TurnId};
247
248        let session = std::sync::Arc::new(crate::session::Session::open_ephemeral());
249        let run_id = FlowRunId::now();
250        let (stream_tx, mut stream_rx) = tokio::sync::broadcast::channel(8);
251        let messages = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
252        let ctx = ToolCtx::new()
253            .with_anchors(Some(TurnId::now()), Some(run_id.clone()), None)
254            .with_history_segment(crate::tool::HistorySegment::Spawned)
255            .with_events(session.sink().clone())
256            .with_session_messages_handle(messages)
257            .with_stream_tx(stream_tx);
258        let message = Message::context_record(
259            TurnId::now(),
260            ContextRecord::new(
261                "handoff.parent",
262                1,
263                ContextRecordAuthority::Runtime,
264                ContextRecordRetention::Latest,
265                ContextRecordBody::text("delegated"),
266            ),
267        );
268
269        append_message_to_context(&ctx, message).unwrap();
270
271        assert!(matches!(
272            stream_rx.try_recv(),
273            Err(tokio::sync::broadcast::error::TryRecvError::Empty)
274        ));
275        assert!(session.sink().snapshot().iter().any(|event| {
276            matches!(
277                event,
278                Event::SystemMsg {
279                    flow_run_id: Some(owner),
280                    ..
281                } if owner == &run_id
282            )
283        }));
284    }
285}