Skip to main content

lash_core/tool_dispatch/
context.rs

1use std::sync::{Arc, Mutex};
2
3use tokio::sync::mpsc;
4
5use crate::plugin::{
6    PluginSession, SessionGraphService, SessionLifecycleService, SessionStateService,
7};
8use crate::{
9    PreparedToolCall, SessionStreamEvent, ToolCallRecord, ToolCatalog, ToolFailure,
10    ToolFailureClass, ToolProvider, ToolResult,
11};
12
13#[derive(Clone, Default)]
14pub(crate) struct CheckpointMessageBuffer {
15    queue: Arc<Mutex<Vec<crate::PluginMessage>>>,
16}
17
18impl CheckpointMessageBuffer {
19    pub(crate) fn enqueue(&self, messages: Vec<crate::PluginMessage>) -> Result<(), String> {
20        let mut queue = self
21            .queue
22            .lock()
23            .map_err(|_| "checkpoint message buffer poisoned".to_string())?;
24        queue.extend(messages);
25        Ok(())
26    }
27
28    pub(crate) fn drain(&self) -> Result<Vec<crate::PluginMessage>, String> {
29        let mut queue = self
30            .queue
31            .lock()
32            .map_err(|_| "checkpoint message buffer poisoned".to_string())?;
33        Ok(queue.drain(..).collect())
34    }
35}
36
37#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
38pub struct ToolTriggerEffectOutcome {
39    pub source_type: String,
40    pub source_key: String,
41    pub occurrence_id: String,
42    #[serde(default)]
43    pub payload: serde_json::Value,
44    pub idempotency_key: String,
45    #[serde(default, skip_serializing_if = "Option::is_none")]
46    pub source: Option<serde_json::Value>,
47    pub deliveries: Vec<crate::TriggerDeliveryEmitReport>,
48}
49
50#[derive(Clone, Default)]
51pub(crate) struct ToolTriggerOutcomeBuffer {
52    queue: Arc<Mutex<Vec<ToolTriggerEffectOutcome>>>,
53}
54
55impl ToolTriggerOutcomeBuffer {
56    pub(crate) fn enqueue(&self, outcome: ToolTriggerEffectOutcome) -> Result<(), String> {
57        let mut queue = self
58            .queue
59            .lock()
60            .map_err(|_| "tool trigger outcome buffer poisoned".to_string())?;
61        queue.push(outcome);
62        Ok(())
63    }
64
65    pub(crate) fn drain(&self) -> Result<Vec<ToolTriggerEffectOutcome>, String> {
66        let mut queue = self
67            .queue
68            .lock()
69            .map_err(|_| "tool trigger outcome buffer poisoned".to_string())?;
70        Ok(queue.drain(..).collect())
71    }
72}
73
74#[derive(Clone)]
75pub struct ToolDispatchContext<'run> {
76    pub plugins: Arc<PluginSession>,
77    pub tools: Arc<dyn ToolProvider>,
78    pub tool_catalog: Arc<ToolCatalog>,
79    pub sessions: Arc<dyn SessionStateService>,
80    pub session_lifecycle: Arc<dyn SessionLifecycleService>,
81    pub session_graph: Arc<dyn SessionGraphService>,
82    pub processes: Arc<dyn crate::ProcessService>,
83    pub process_cancel_ability: Arc<dyn crate::ProcessCancelAbility>,
84    pub trigger_router: Option<crate::TriggerRouter>,
85    pub(crate) effect_controller: crate::runtime::RuntimeEffectControllerHandle<'run>,
86    pub(crate) direct_completions: crate::DirectCompletionClient<'run>,
87    pub(crate) parent_invocation: Option<crate::RuntimeInvocation>,
88    pub(crate) execution_env_spec: crate::ProcessExecutionEnvSpec,
89    pub session_id: String,
90    pub agent_frame_id: crate::AgentFrameId,
91    pub event_tx: mpsc::Sender<SessionStreamEvent>,
92    pub(crate) checkpoint_messages: CheckpointMessageBuffer,
93    pub(crate) trigger_outcomes: ToolTriggerOutcomeBuffer,
94    pub attachment_store: Arc<crate::SessionAttachmentStore>,
95    pub attachment_source_policy: Arc<dyn crate::AttachmentSourcePolicy>,
96    pub turn_context: crate::TurnContext,
97    pub clock: Arc<dyn crate::Clock>,
98}
99
100impl<'run> ToolDispatchContext<'run> {
101    pub fn process_scope(&self) -> crate::ProcessOpScope<'_> {
102        crate::ProcessOpScope::new(self.effect_controller.scoped())
103            .with_parent_invocation(self.parent_invocation.clone())
104            .with_agent_frame_id(Some(self.agent_frame_id.clone()))
105    }
106
107    pub(crate) fn to_static(&self) -> Option<ToolDispatchContext<'static>> {
108        Some(ToolDispatchContext {
109            plugins: Arc::clone(&self.plugins),
110            tools: Arc::clone(&self.tools),
111            tool_catalog: Arc::clone(&self.tool_catalog),
112            sessions: Arc::clone(&self.sessions),
113            session_lifecycle: Arc::clone(&self.session_lifecycle),
114            session_graph: Arc::clone(&self.session_graph),
115            processes: Arc::clone(&self.processes),
116            process_cancel_ability: Arc::clone(&self.process_cancel_ability),
117            trigger_router: self.trigger_router.clone(),
118            effect_controller: self.effect_controller.to_static()?,
119            direct_completions: self.direct_completions.to_static()?,
120            parent_invocation: self.parent_invocation.clone(),
121            execution_env_spec: self.execution_env_spec.clone(),
122            session_id: self.session_id.clone(),
123            agent_frame_id: self.agent_frame_id.clone(),
124            event_tx: self.event_tx.clone(),
125            checkpoint_messages: self.checkpoint_messages.clone(),
126            trigger_outcomes: self.trigger_outcomes.clone(),
127            attachment_store: Arc::clone(&self.attachment_store),
128            attachment_source_policy: Arc::clone(&self.attachment_source_policy),
129            turn_context: self.turn_context.clone(),
130            clock: Arc::clone(&self.clock),
131        })
132    }
133}
134
135#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
136pub(crate) struct ToolDispatchOutcome {
137    pub record: ToolCallRecord,
138}
139
140#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
141pub(crate) struct PendingToolDispatchOutcome {
142    pub tool_name: String,
143    pub args: serde_json::Value,
144    pub key: crate::AwaitEventKey,
145    pub pending: crate::PendingCompletion,
146    pub duration_ms: u64,
147}
148
149#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
150#[serde(tag = "status", rename_all = "snake_case")]
151pub(crate) enum ToolCallLaunch {
152    Done(ToolDispatchOutcome),
153    Pending(PendingToolDispatchOutcome),
154}
155
156pub(crate) enum ToolPreparationOutcome {
157    Prepared(PreparedToolCall),
158    Completed(Box<ToolDispatchOutcome>),
159}
160
161pub(super) fn completed_preparation(outcome: ToolDispatchOutcome) -> ToolPreparationOutcome {
162    ToolPreparationOutcome::Completed(Box::new(outcome))
163}
164pub(super) fn outcome(
165    tool_name: String,
166    args: serde_json::Value,
167    result: ToolResult,
168    duration_ms: u64,
169) -> ToolDispatchOutcome {
170    let record = ToolCallRecord {
171        call_id: None,
172        tool: tool_name,
173        args,
174        output: result.into_done_output().unwrap_or_else(|_| {
175            crate::ToolCallOutput::failure(crate::ToolFailure::runtime(
176                crate::ToolFailureClass::Internal,
177                "pending_tool_not_finalized",
178                "pending tool result reached a completed-output projection path",
179            ))
180        }),
181        duration_ms,
182    };
183    ToolDispatchOutcome { record }
184}
185
186pub(super) fn launch_done(outcome: ToolDispatchOutcome) -> ToolCallLaunch {
187    ToolCallLaunch::Done(outcome)
188}
189
190pub(super) fn runtime_failure(
191    class: ToolFailureClass,
192    code: impl Into<String>,
193    message: impl Into<String>,
194) -> ToolResult {
195    ToolResult::failure(ToolFailure::runtime(class, code, message))
196}