Skip to main content

mobius/middleware/
context.rs

1use std::borrow::Cow;
2use std::collections::{BTreeMap, BTreeSet};
3use std::sync::Arc;
4
5use serde_json::Value;
6
7use super::MiddlewareStack;
8use super::tools::Catalog;
9use super::tools::ToolResult;
10use super::{approximate_item_tokens, approximate_tokens, serialized_len};
11use crate::agent::{AgentRole, WeakAgentSender};
12use crate::backend::checkpoint::{
13    Checkpoint, CheckpointStore, ContextRewriteReason, ExecutionOutcome, MAX_QUEUED_MESSAGES,
14    QueuedMessage as DurableQueuedMessage, QueuedMessageBoundary,
15};
16use crate::backend::model::{ModelRouter, message_input};
17use crate::backend::sandbox::ApprovalPolicy;
18use crate::protocol::{
19    EventMsg, FrontendEvent, MAX_CAPABILITY_INPUT_BYTES, MessageAuthor, MessageEvent,
20    MessageSubmission, MessageTarget, ReviewDecision, SessionContext, SessionFileReference,
21    TokenUsage, ToolCall, message_metadata,
22};
23use crate::{Error, Result};
24
25/// Sends middleware-owned UI updates without depending on a concrete frontend.
26pub type FrontendEventSink = Arc<dyn Fn(FrontendEvent) -> Result<()> + Send + Sync>;
27
28/// Read-only queued message owned by the middleware receiving it.
29#[derive(Debug, Clone, Copy, PartialEq)]
30pub struct QueuedMessageView<'a> {
31    item: &'a DurableQueuedMessage,
32}
33
34impl<'a> QueuedMessageView<'a> {
35    /// Returns the identity token required by a conditional queue mutation.
36    #[must_use]
37    pub fn id(&self) -> &'a str {
38        self.item.id()
39    }
40
41    /// Returns the prepared presentation event.
42    #[must_use]
43    pub fn event(&self) -> MessageEvent {
44        self.item.event()
45    }
46}
47
48/// Read-only startup snapshot containing only one middleware's queued messages.
49#[derive(Clone, Default)]
50pub struct QueuedMessageSnapshot {
51    items: Vec<DurableQueuedMessage>,
52}
53
54impl QueuedMessageSnapshot {
55    /// Returns every queued item owned by this middleware, oldest first.
56    pub fn views(&self) -> impl Iterator<Item = QueuedMessageView<'_>> {
57        self.items.iter().map(|item| QueuedMessageView { item })
58    }
59
60    pub(super) fn for_owner(owner: &str, items: &[DurableQueuedMessage]) -> Self {
61        Self {
62            items: items
63                .iter()
64                .filter(|item| item.owner() == owner)
65                .cloned()
66                .collect(),
67        }
68    }
69}
70
71/// Mutable scoped view of messages retained until their delivery boundary.
72pub struct MessageQueue<'a> {
73    items: &'a mut Vec<DurableQueuedMessage>,
74    owner: Option<&'static str>,
75}
76
77impl<'a> MessageQueue<'a> {
78    pub(crate) fn new(items: &'a mut Vec<DurableQueuedMessage>) -> Self {
79        Self { items, owner: None }
80    }
81
82    pub(super) fn scope(&mut self, owner: &'static str) {
83        self.owner = Some(owner);
84    }
85
86    fn owner(&self) -> Result<&'static str> {
87        self.owner
88            .ok_or_else(|| Error::Config("message queue is not scoped to a middleware".into()))
89    }
90
91    /// Returns the number of queued items owned by this middleware.
92    #[must_use]
93    pub fn count(&self) -> usize {
94        let Some(owner) = self.owner else {
95            return 0;
96        };
97        self.items
98            .iter()
99            .filter(|item| item.owner() == owner)
100            .count()
101    }
102
103    /// Returns the newest message available to this context.
104    #[must_use]
105    pub fn latest(&self) -> Option<QueuedMessageView<'_>> {
106        let owner = self.owner?;
107        self.items
108            .iter()
109            .rev()
110            .find(|item| item.owner() == owner)
111            .map(|item| QueuedMessageView { item })
112    }
113
114    /// Returns one owned item by its revision identity.
115    #[must_use]
116    pub fn find(&self, id: &str) -> Option<QueuedMessageView<'_>> {
117        let owner = self.owner?;
118        self.items
119            .iter()
120            .find(|item| item.owner() == owner && item.id() == id)
121            .map(|item| QueuedMessageView { item })
122    }
123
124    /// Appends one prepared message, or returns `false` when it is full or duplicated.
125    /// # Errors
126    ///
127    /// Returns an error if validation or an operation required by this function fails.
128    pub fn enqueue(
129        &mut self,
130        id: &str,
131        boundary: QueuedMessageBoundary,
132        event: MessageEvent,
133    ) -> Result<bool> {
134        let owner = self.owner()?;
135        let item = DurableQueuedMessage::new(owner, id, boundary, event)?;
136        if self.items.len() >= MAX_QUEUED_MESSAGES {
137            return Ok(false);
138        }
139        if self
140            .items
141            .iter()
142            .any(|item| item.owner() == owner && item.id() == id)
143        {
144            return Ok(false);
145        }
146        self.items.push(item);
147        Ok(true)
148    }
149
150    /// Atomically replaces one owned item while preserving its queue position.
151    /// # Errors
152    ///
153    /// Returns an error if validation or an operation required by this function fails.
154    pub fn replace(&mut self, id: &str, replacement_id: &str, event: MessageEvent) -> Result<bool> {
155        let owner = self.owner()?;
156        let Some(index) = self
157            .items
158            .iter()
159            .position(|item| item.owner() == owner && item.id() == id)
160        else {
161            return Ok(false);
162        };
163        if self.items.iter().enumerate().any(|(candidate, item)| {
164            candidate != index && item.owner() == owner && item.id() == replacement_id
165        }) {
166            return Ok(false);
167        }
168        self.items[index].replace(replacement_id, event)?;
169        Ok(true)
170    }
171
172    pub(crate) fn stage_model_messages(&mut self, turn_id: &str) -> Result<Vec<PreparedMessage>> {
173        let Some(owner) = self.owner else {
174            return Ok(Vec::new());
175        };
176        self.items
177            .extract_if(.., |item| {
178                item.owner() == owner
179                    && matches!(
180                        item.boundary(),
181                        QueuedMessageBoundary::Steer { turn_id: target }
182                            if target == turn_id
183                    )
184            })
185            .map(PreparedMessage::try_from)
186            .collect()
187    }
188
189    pub(crate) fn next_turn(&self) -> Result<Option<PreparedMessage>> {
190        let owner = self.owner()?;
191        self.items
192            .iter()
193            .find(|item| item.owner() == owner && item.boundary().starts_turn())
194            .cloned()
195            .map(PreparedMessage::try_from)
196            .transpose()
197    }
198
199    pub(crate) fn consume_next_turn(&mut self, id: &str) -> Result<()> {
200        let owner = self.owner()?;
201        let index = self
202            .items
203            .iter()
204            .position(|item| {
205                item.owner() == owner && item.id() == id && item.boundary().starts_turn()
206            })
207            .ok_or_else(|| Error::Checkpoint("prepared message is no longer queued".into()))?;
208        self.items.remove(index);
209        Ok(())
210    }
211
212    pub(crate) fn promote_failed_turn(&mut self, turn_id: &str) -> Result<()> {
213        let owner = self.owner()?;
214        for item in self.items.iter_mut().filter(|item| {
215            item.owner() == owner
216                && matches!(
217                    item.boundary(),
218                    QueuedMessageBoundary::Steer { turn_id: target }
219                        if target == turn_id
220                )
221        }) {
222            item.promote_to_next_turn()?;
223        }
224        Ok(())
225    }
226}
227
228/// One queued message prepared for its model boundary.
229pub(crate) struct PreparedMessage {
230    pub(crate) submission_id: String,
231    pub(crate) input: Value,
232    pub(crate) event: EventMsg,
233    pub(crate) title_seed: Option<String>,
234    pub(crate) boundary_events: Vec<EventMsg>,
235}
236
237impl TryFrom<DurableQueuedMessage> for PreparedMessage {
238    type Error = Error;
239
240    fn try_from(message: DurableQueuedMessage) -> Result<Self> {
241        let (submission_id, event) = message.into_parts();
242        let input = message_input(&event)?;
243        let title_seed = matches!(
244            event.author,
245            MessageAuthor::User | MessageAuthor::Source { .. }
246        )
247        .then(|| event.text.trim().to_string())
248        .filter(|title| !title.is_empty());
249        Ok(Self {
250            submission_id,
251            input,
252            event: EventMsg::Message(event),
253            title_seed,
254            boundary_events: Vec::new(),
255        })
256    }
257}
258
259/// Durable runtime identity exposed while middleware starts a session.
260#[derive(Clone)]
261pub struct RuntimeContext {
262    /// The sender.
263    pub sender: WeakAgentSender,
264    /// The checkpoints.
265    pub checkpoints: Arc<dyn CheckpointStore>,
266    /// The session identifier.
267    pub session_id: String,
268    /// The model route.
269    pub model_route: String,
270    /// The model.
271    pub model: String,
272    /// The approval policy.
273    pub approval_policy: ApprovalPolicy,
274    /// The session context.
275    pub session_context: SessionContext,
276    /// The metadata.
277    pub metadata: BTreeMap<String, Value>,
278    /// The role.
279    pub role: AgentRole,
280    /// The frontend.
281    pub frontend: FrontendEventSink,
282}
283
284impl RuntimeContext {
285    pub(crate) fn turn_identity<'a>(
286        &'a self,
287        turn_id: &'a str,
288        author: &'a MessageAuthor,
289    ) -> TurnIdentity<'a> {
290        TurnIdentity {
291            session_id: &self.session_id,
292            turn_id,
293            model: &self.model,
294            approval_policy: self.approval_policy,
295            author,
296        }
297    }
298}
299
300/// Stable facts shared by hooks that run within one active turn.
301#[derive(Debug, Clone, Copy, PartialEq, Eq)]
302pub struct TurnIdentity<'a> {
303    /// Trusted provenance of the message that initiated the turn.
304    pub author: &'a MessageAuthor,
305    /// The session identifier.
306    pub session_id: &'a str,
307    /// The turn identifier.
308    pub turn_id: &'a str,
309    /// The model.
310    pub model: &'a str,
311    /// The approval policy.
312    pub approval_policy: ApprovalPolicy,
313}
314
315/// Why [`Middleware::session_start`](super::Middleware::session_start) is running.
316#[derive(Debug, Clone, Copy, PartialEq, Eq)]
317pub enum SessionStartSource {
318    /// Selects the startup case.
319    Startup,
320    /// Selects the resume case.
321    Resume,
322    /// Selects the compact case.
323    Compact,
324}
325
326/// Mutable state shared by the declaration-ordered `SessionStart` hooks.
327pub struct SessionStartContext<'a> {
328    /// The runtime.
329    pub runtime: &'a RuntimeContext,
330    pub(crate) source: SessionStartSource,
331    pub(crate) queued_messages: QueuedMessageSnapshot,
332    pub(crate) input: &'a mut Vec<Value>,
333    pub(crate) input_changed: bool,
334    pub(crate) stop_reason: Option<String>,
335}
336
337impl SessionStartContext<'_> {
338    #[must_use]
339    /// Returns the session-start source.
340    pub fn source(&self) -> SessionStartSource {
341        self.source
342    }
343
344    #[must_use]
345    /// Returns the queued messages.
346    pub fn queued_messages(&self) -> &QueuedMessageSnapshot {
347        &self.queued_messages
348    }
349
350    /// Appends hidden provider context produced while the session starts.
351    pub fn push_input(&mut self, item: Value) {
352        self.input.push(item);
353        self.input_changed = true;
354    }
355
356    pub(crate) fn retain_input(&mut self, mut keep: impl FnMut(&Value) -> bool) {
357        let input_len = self.input.len();
358        self.input.retain(&mut keep);
359        self.input_changed |= self.input.len() != input_len;
360    }
361
362    /// Stops the active turn after session-start processing completes.
363    /// # Errors
364    ///
365    /// Returns an error if validation or an operation required by this function fails.
366    pub fn stop(&mut self, reason: impl Into<String>) -> Result<()> {
367        set_stop_reason(&mut self.stop_reason, "session-start stop", reason)
368    }
369
370    /// Returns the first stop requested by the ordered middleware chain.
371    #[must_use]
372    pub fn stop_reason(&self) -> Option<&str> {
373        self.stop_reason.as_deref()
374    }
375}
376
377/// Mutable state exposed before a prepared next-turn message enters durable context.
378pub struct MessageSubmitContext<'a> {
379    /// The turn.
380    pub turn: TurnIdentity<'a>,
381    /// The author.
382    pub author: &'a MessageAuthor,
383    /// The message.
384    pub message: &'a str,
385    /// The attachments.
386    pub attachments: &'a [SessionFileReference],
387    /// The events.
388    pub events: &'a mut Vec<EventMsg>,
389    pub(crate) input: Vec<Value>,
390    pub(crate) rejection: Option<String>,
391}
392
393impl MessageSubmitContext<'_> {
394    /// Adds provider-neutral context immediately before the submitted message.
395    pub fn push_input(&mut self, item: Value) {
396        self.input.push(item);
397    }
398
399    /// Rejects the submission without treating the policy decision as a hook failure.
400    /// # Errors
401    ///
402    /// Returns an error if validation or an operation required by this function fails.
403    pub fn reject(&mut self, reason: impl Into<String>) -> Result<()> {
404        let reason = hook_message("prompt rejection", reason)?;
405        if self.rejection.is_none() {
406            self.rejection = Some(reason);
407        }
408        Ok(())
409    }
410}
411
412pub(crate) struct MessageSubmitResult {
413    pub(crate) input: Vec<Value>,
414    pub(crate) rejection: Option<String>,
415}
416
417/// Mutable state exposed immediately before a model request.
418pub struct ModelContext<'a> {
419    /// Trusted provenance of the message that initiated the active turn.
420    pub author: &'a MessageAuthor,
421    /// The model.
422    pub model: &'a ModelRouter,
423    /// The provider.
424    pub provider: &'a str,
425    /// The session identifier.
426    pub session_id: &'a str,
427    /// The session context.
428    pub session_context: &'a SessionContext,
429    /// The metadata.
430    pub metadata: &'a BTreeMap<String, Value>,
431    /// The turn identifier.
432    pub turn_id: &'a str,
433    /// The model step.
434    pub model_step: usize,
435    /// The context window.
436    pub context_window: i64,
437    /// The instructions.
438    pub instructions: &'a str,
439    pub(crate) checkpoint_sequence: u64,
440    pub(crate) available_tools: &'a mut BTreeSet<String>,
441    pub(crate) allow_hosted_tools: &'a mut bool,
442    pub(crate) durable_input: &'a mut Vec<Value>,
443    pub(crate) transcript_delta: &'a mut Vec<Value>,
444    pub(crate) context_epoch: &'a mut u64,
445    pub(crate) compaction_count: &'a mut u64,
446    pub(crate) rewrite_reasons: &'a mut Vec<ContextRewriteReason>,
447    pub(crate) turn_stop: &'a mut Option<String>,
448    pub(crate) queued_messages: Vec<DurableQueuedMessage>,
449    /// The last usage.
450    pub last_usage: Option<&'a TokenUsage>,
451    /// The tools.
452    pub tools: &'a Catalog,
453    /// The events.
454    pub events: &'a mut Vec<EventMsg>,
455    /// The usage.
456    pub usage: &'a mut Vec<TokenUsage>,
457    /// Set when this hook changes durable checkpoint state.
458    pub(crate) checkpoint_changed: &'a mut bool,
459    pub(crate) runtime: &'a RuntimeContext,
460    pub(crate) hooks: &'a MiddlewareStack,
461}
462
463/// Live capability state used to hide registered tools at a model boundary.
464pub struct ToolExposureContext<'a> {
465    /// The session identifier.
466    pub session_id: &'a str,
467    pub(crate) supports_tool_image_input: bool,
468    pub(crate) input: &'a [Value],
469    pub(crate) available: &'a mut BTreeSet<String>,
470}
471
472impl ToolExposureContext<'_> {
473    /// Reports whether the active model accepts image input.
474    #[must_use]
475    pub fn supports_tool_image_input(&self) -> bool {
476        self.supports_tool_image_input
477    }
478
479    /// Returns the most recent typed conversation message in model context.
480    #[must_use]
481    pub fn latest_message(&self) -> Option<MessageEvent> {
482        self.input.iter().rev().find_map(message_metadata)
483    }
484
485    /// Hides registered tools for this boundary.
486    pub fn hide(&mut self, names: &[&str]) {
487        for name in names {
488            self.available.remove(*name);
489        }
490    }
491}
492
493impl ModelContext<'_> {
494    /// Prevents provider-hosted tools for this model step.
495    pub fn disable_hosted_tools(&mut self) {
496        *self.allow_hosted_tools = false;
497    }
498
499    /// Returns durable provider-neutral model context.
500    #[must_use]
501    pub fn input(&self) -> &[Value] {
502        self.durable_input
503    }
504
505    /// Replaces active model context and advances its rewrite epoch once per boundary.
506    /// # Errors
507    ///
508    /// Returns an error if validation or an operation required by this function fails.
509    pub fn rewrite_input(
510        &mut self,
511        reason: ContextRewriteReason,
512        mut input: Vec<Value>,
513    ) -> Result<()> {
514        if *self.durable_input == input {
515            return Ok(());
516        }
517        if self.rewrite_reasons.is_empty() {
518            *self.context_epoch = self
519                .context_epoch
520                .checked_add(1)
521                .ok_or_else(|| Error::Checkpoint("context rewrite epoch overflow".into()))?;
522        }
523        if !self.rewrite_reasons.contains(&reason) {
524            self.rewrite_reasons.push(reason);
525        }
526        crate::backend::model::reset_prompt_cache_breakpoint(&mut input);
527        *self.durable_input = input;
528        self.last_usage = None;
529        *self.checkpoint_changed = true;
530        Ok(())
531    }
532
533    /// Appends a durable replay item without adding it to provider context.
534    pub(crate) fn record_transcript_item(&mut self, item: Value) {
535        self.transcript_delta.push(item);
536        *self.checkpoint_changed = true;
537    }
538
539    /// Appends durable provider context without adding synthetic replay history.
540    pub fn append_model_input(&mut self, item: Value) {
541        self.durable_input.push(item);
542        *self.checkpoint_changed = true;
543    }
544
545    /// Appends durable input to model context and its transcript journal.
546    /// # Errors
547    ///
548    /// Returns an error if validation or an operation required by this function fails.
549    pub fn push_input(&mut self, item: Value) -> Result<MessageTarget> {
550        self.durable_input.push(item.clone());
551        self.transcript_delta.push(item);
552        *self.checkpoint_changed = true;
553        provisional_message_target(self.checkpoint_sequence, self.transcript_delta.len())
554    }
555
556    /// Estimates visible history, instructions, and tool schemas at four bytes per token.
557    #[must_use]
558    pub fn estimated_input_tokens(&self) -> i64 {
559        let Ok(tools) = self
560            .tools
561            .prepare(self.input(), self.available_tools.clone())
562        else {
563            return i64::MAX;
564        };
565        let visible = tools
566            .direct()
567            .iter()
568            .chain(
569                tools
570                    .deferred()
571                    .iter()
572                    .filter(|tool| tools.materialized().contains(&tool.name)),
573            )
574            .collect::<Vec<_>>();
575        let Some(tool_bytes) = serialized_len(&visible) else {
576            return i64::MAX;
577        };
578        let history = self
579            .durable_input
580            .iter()
581            .map(approximate_item_tokens)
582            .fold(0usize, usize::saturating_add);
583        i64::try_from(history.saturating_add(approximate_tokens(
584            tool_bytes.saturating_add(self.instructions.len()),
585        )))
586        .unwrap_or(i64::MAX)
587    }
588
589    pub(crate) async fn pre_compact(&mut self) -> Result<()> {
590        let hooks = self.hooks;
591        let stop_reason = hooks
592            .pre_compact(CompactContext {
593                session_id: self.session_id,
594                turn_id: self.turn_id,
595                model: &self.runtime.model,
596                input: self.durable_input,
597                events: self.events,
598                stop_reason: None,
599            })
600            .await?;
601        set_first(self.turn_stop, stop_reason);
602        Ok(())
603    }
604
605    pub(crate) async fn post_compact(&mut self) -> Result<()> {
606        let hooks = self.hooks;
607        let stop_reason = hooks
608            .post_compact(CompactContext {
609                session_id: self.session_id,
610                turn_id: self.turn_id,
611                model: &self.runtime.model,
612                input: self.durable_input,
613                events: self.events,
614                stop_reason: None,
615            })
616            .await?;
617        set_first(self.turn_stop, stop_reason);
618        if self.turn_stop.is_some() {
619            return Ok(());
620        }
621        let start = hooks
622            .session_start(
623                self.runtime,
624                &self.queued_messages,
625                SessionStartSource::Compact,
626                self.durable_input,
627            )
628            .await?;
629        set_first(self.turn_stop, start.stop_reason);
630        Ok(())
631    }
632
633    #[must_use]
634    pub(crate) fn turn_stopped(&self) -> bool {
635        self.turn_stop.is_some()
636    }
637}
638
639/// Request-only model input exposed after every durable `PreModel` hook.
640pub struct ModelRequestContext<'a> {
641    /// Trusted provenance of the message that initiated the active turn.
642    pub author: &'a MessageAuthor,
643    /// The role.
644    pub role: &'a AgentRole,
645    /// The model.
646    pub model: &'a ModelRouter,
647    /// The provider.
648    pub provider: &'a str,
649    /// The session identifier.
650    pub session_id: &'a str,
651    /// The turn identifier.
652    pub turn_id: &'a str,
653    /// The model step.
654    pub model_step: usize,
655    pub(crate) input: Cow<'a, [Value]>,
656}
657
658impl ModelRequestContext<'_> {
659    /// Returns the input currently prepared for this one model request.
660    #[must_use]
661    pub fn input(&self) -> &[Value] {
662        self.input.as_ref()
663    }
664
665    /// Replaces only the input sent by this model request.
666    pub fn replace_input(&mut self, input: Vec<Value>) {
667        self.input = Cow::Owned(input);
668    }
669}
670
671/// Mutable policy boundary for one normalized model-requested tool call.
672pub struct PreToolUseContext<'a> {
673    /// The turn.
674    pub turn: TurnIdentity<'a>,
675    /// The events.
676    pub events: &'a mut Vec<EventMsg>,
677    pub(crate) tools: &'a Catalog,
678    pub(crate) call: &'a mut ToolCall,
679    pub(crate) input: Vec<Value>,
680    pub(crate) denial: Option<String>,
681}
682
683impl PreToolUseContext<'_> {
684    /// Returns the call after any earlier middleware rewrites.
685    #[must_use]
686    pub fn call(&self) -> &ToolCall {
687        self.call
688    }
689
690    /// Replaces the tool name and arguments while preserving the provider call ID.
691    /// # Errors
692    ///
693    /// Returns an error if validation or an operation required by this function fails.
694    pub fn replace(&mut self, name: impl Into<String>, arguments: Value) -> Result<()> {
695        self.call.replace(name.into(), arguments)
696    }
697
698    /// Adds durable provider-neutral context before this call at a tool-complete boundary.
699    pub fn push_input(&mut self, item: Value) {
700        self.input.push(item);
701    }
702
703    /// Denies the call. Later middleware may observe but cannot undo the denial.
704    /// # Errors
705    ///
706    /// Returns an error if validation or an operation required by this function fails.
707    pub fn deny(&mut self, reason: impl Into<String>) -> Result<()> {
708        let reason = hook_message("tool denial", reason)?;
709        if self.denial.is_none() {
710            self.denial = Some(reason);
711        }
712        Ok(())
713    }
714
715    /// Returns the first denial made by the ordered middleware chain.
716    #[must_use]
717    pub fn denial(&self) -> Option<&str> {
718        self.denial.as_deref()
719    }
720}
721
722/// Mutable policy boundary for a sandbox approval request.
723pub struct PermissionRequestContext<'a> {
724    /// The turn.
725    pub turn: TurnIdentity<'a>,
726    /// The calls.
727    pub calls: &'a [ToolCall],
728    /// The requested call identifiers.
729    pub requested_call_ids: &'a [String],
730    /// The reason.
731    pub reason: &'a str,
732    /// The events.
733    pub events: &'a mut Vec<EventMsg>,
734    pub(crate) tools: &'a Catalog,
735    pub(crate) decision: Option<ReviewDecision>,
736}
737
738impl PermissionRequestContext<'_> {
739    /// Returns the decision accumulated from earlier middleware.
740    #[must_use]
741    pub fn decision(&self) -> Option<&ReviewDecision> {
742        self.decision.as_ref()
743    }
744
745    /// Allows this request unless an earlier middleware denied it.
746    pub fn allow(&mut self) {
747        if !matches!(self.decision, Some(ReviewDecision::Denied { .. })) {
748            self.decision = Some(ReviewDecision::Approved);
749        }
750    }
751
752    /// Denies this request. The decision cannot be weakened by later middleware.
753    /// # Errors
754    ///
755    /// Returns an error if validation or an operation required by this function fails.
756    pub fn deny(&mut self, reason: impl Into<String>) -> Result<()> {
757        let reason = hook_message("permission denial", reason)?;
758        if !matches!(self.decision, Some(ReviewDecision::Denied { .. })) {
759            self.decision = Some(ReviewDecision::Denied { rejection: reason });
760        }
761        Ok(())
762    }
763}
764
765/// Mutable model-visible result exposed after an executed tool call.
766pub struct PostToolUseContext<'a> {
767    /// The turn.
768    pub turn: TurnIdentity<'a>,
769    /// The call.
770    pub call: &'a ToolCall,
771    /// The events.
772    pub events: &'a mut Vec<EventMsg>,
773    pub(crate) tools: &'a Catalog,
774    pub(crate) result: &'a mut ToolResult,
775}
776
777impl PostToolUseContext<'_> {
778    /// Returns the result after any earlier middleware changes.
779    #[must_use]
780    pub fn result(&self) -> &ToolResult {
781        self.result
782    }
783
784    /// Replaces the feedback returned to the model without changing past side effects.
785    pub fn replace(&mut self, output: impl Into<String>) {
786        self.result.replace(output.into());
787    }
788
789    /// Adds provider-neutral context immediately after this tool output.
790    pub fn push_input(&mut self, item: Value) {
791        self.result.additional_input.push(item);
792    }
793}
794
795/// State exposed immediately before or after context compaction.
796pub struct CompactContext<'a> {
797    /// The session identifier.
798    pub session_id: &'a str,
799    /// The turn identifier.
800    pub turn_id: &'a str,
801    /// The model.
802    pub model: &'a str,
803    /// The input.
804    pub input: &'a [Value],
805    /// The events.
806    pub events: &'a mut Vec<EventMsg>,
807    pub(crate) stop_reason: Option<String>,
808}
809
810impl CompactContext<'_> {
811    /// Stops the active turn at this compaction boundary.
812    /// # Errors
813    ///
814    /// Returns an error if validation or an operation required by this function fails.
815    pub fn stop(&mut self, reason: impl Into<String>) -> Result<()> {
816        set_stop_reason(&mut self.stop_reason, "compaction stop", reason)
817    }
818
819    /// Returns the first stop requested by the ordered middleware chain.
820    #[must_use]
821    pub fn stop_reason(&self) -> Option<&str> {
822        self.stop_reason.as_deref()
823    }
824}
825
826/// Mutable policy boundary immediately before normal turn completion.
827pub struct StopContext<'a> {
828    /// The turn.
829    pub turn: TurnIdentity<'a>,
830    /// The events.
831    pub events: &'a mut Vec<EventMsg>,
832    pub(crate) role: &'a AgentRole,
833    pub(crate) stop_hook_active: bool,
834    pub(crate) last_assistant_message: Option<&'a str>,
835    pub(crate) continuation: Option<String>,
836}
837
838impl StopContext<'_> {
839    #[must_use]
840    /// Returns the agent role.
841    pub fn role(&self) -> &AgentRole {
842        self.role
843    }
844
845    #[must_use]
846    /// Stops hook active.
847    pub fn stop_hook_active(&self) -> bool {
848        self.stop_hook_active
849    }
850
851    #[must_use]
852    /// Returns the last assistant message.
853    pub fn last_assistant_message(&self) -> Option<&str> {
854        self.last_assistant_message
855    }
856
857    /// Returns the first continuation requested by the middleware chain.
858    #[must_use]
859    pub fn continuation(&self) -> Option<&str> {
860        self.continuation.as_deref()
861    }
862
863    /// Requests one more model step with hidden context.
864    /// # Errors
865    ///
866    /// Returns an error if validation or an operation required by this function fails.
867    pub fn continue_with(&mut self, prompt: impl Into<String>) -> Result<()> {
868        if self.stop_hook_active {
869            return Err(Error::Config(
870                "a stop hook may continue a turn only once".into(),
871            ));
872        }
873        let prompt = hook_message("stop continuation prompt", prompt)?;
874        if self.continuation.is_none() {
875            self.continuation = Some(prompt);
876        }
877        Ok(())
878    }
879}
880
881fn hook_message(name: &str, value: impl Into<String>) -> Result<String> {
882    let value = value.into();
883    if value.trim().is_empty() || value.len() > MAX_CAPABILITY_INPUT_BYTES {
884        return Err(Error::Config(format!("{name} is empty or too long")));
885    }
886    Ok(value)
887}
888
889fn set_stop_reason(
890    target: &mut Option<String>,
891    name: &str,
892    reason: impl Into<String>,
893) -> Result<()> {
894    let reason = hook_message(name, reason)?;
895    if target.is_none() {
896        *target = Some(reason);
897    }
898    Ok(())
899}
900
901fn set_first(target: &mut Option<String>, value: Option<String>) {
902    if target.is_none() {
903        *target = value;
904    }
905}
906
907pub(super) fn provisional_message_target(
908    checkpoint_sequence: u64,
909    batch_item_count: usize,
910) -> Result<MessageTarget> {
911    Ok(MessageTarget {
912        checkpoint_sequence: checkpoint_sequence
913            .checked_add(1)
914            .ok_or_else(|| Error::Checkpoint("checkpoint sequence overflow".into()))?,
915        batch_item_count,
916    })
917}
918
919/// Mutable state exposed to the middleware preparing conversation messages.
920pub struct MessageRouteContext<'a> {
921    /// The submission identifier.
922    pub submission_id: &'a str,
923    /// The message.
924    pub message: &'a MessageSubmission,
925    /// The active turn identifier.
926    pub active_turn_id: Option<&'a str>,
927    /// The queued messages.
928    pub queued_messages: MessageQueue<'a>,
929    /// The events.
930    pub events: &'a mut Vec<EventMsg>,
931}
932
933/// Mutable turn state exposed to a capability command that can run immediately.
934pub struct ActiveCommandContext<'a> {
935    /// The checkpoints.
936    pub checkpoints: &'a dyn CheckpointStore,
937    /// The submission identifier.
938    pub submission_id: &'a str,
939    /// The session identifier.
940    pub session_id: &'a str,
941    /// The metadata.
942    pub metadata: &'a BTreeMap<String, Value>,
943    /// The active turn identifier.
944    pub active_turn_id: &'a str,
945    /// The command.
946    pub command: &'a str,
947    /// The arguments.
948    pub arguments: &'a str,
949    /// The input.
950    pub input: Option<&'a str>,
951    /// The target.
952    pub target: Option<MessageTarget>,
953    /// The queued messages.
954    pub queued_messages: MessageQueue<'a>,
955    /// The events.
956    pub events: &'a mut Vec<EventMsg>,
957}
958
959/// Result of one middleware-owned submission.
960#[derive(Debug, Clone, PartialEq, Eq)]
961pub enum SubmissionResult {
962    /// Selects the accepted case.
963    Accepted {
964        /// The input changed.
965        input_changed: bool,
966    },
967    /// The operation completed without changing durable turn state; publish its events now.
968    Handled,
969    /// Selects the rejected case.
970    Rejected(String),
971}
972
973/// State exposed when the loop finishes or aborts a turn.
974pub struct TurnEndContext<'a> {
975    /// The session identifier.
976    pub session_id: &'a str,
977    /// The turn identifier.
978    pub turn_id: &'a str,
979    pub(crate) outcome: ExecutionOutcome,
980    pub(crate) queued_messages: &'a [DurableQueuedMessage],
981    pub(crate) owner: Option<&'static str>,
982    /// The events.
983    pub events: &'a mut Vec<EventMsg>,
984}
985
986impl TurnEndContext<'_> {
987    #[must_use]
988    /// Returns the turn outcome.
989    pub fn outcome(&self) -> ExecutionOutcome {
990        self.outcome
991    }
992
993    /// Returns queued messages still pending for this middleware, oldest first.
994    pub fn queued_messages(&self) -> impl Iterator<Item = QueuedMessageView<'_>> {
995        let owner = self.owner;
996        self.queued_messages
997            .iter()
998            .filter(move |item| owner.is_some_and(|owner| item.owner() == owner))
999            .map(|item| QueuedMessageView { item })
1000    }
1001}
1002
1003/// State available to a middleware-owned frontend command.
1004pub struct MiddlewareCommandContext<'a> {
1005    /// The command.
1006    pub command: &'a str,
1007    /// The arguments.
1008    pub arguments: &'a str,
1009    /// The input.
1010    pub input: Option<&'a str>,
1011    /// The target.
1012    pub target: Option<MessageTarget>,
1013    /// The session identifier.
1014    pub session_id: &'a str,
1015    /// The session context.
1016    pub session_context: &'a SessionContext,
1017    /// The checkpoint.
1018    pub checkpoint: &'a Checkpoint,
1019    /// The checkpoints.
1020    pub checkpoints: Arc<dyn CheckpointStore>,
1021}
1022
1023#[cfg(test)]
1024mod tests {
1025    use super::*;
1026    use crate::backend::model::{Model, ModelEventSink, ModelOutput, ModelRequest};
1027
1028    struct NoModel;
1029
1030    impl Model for NoModel {
1031        fn respond<'a>(
1032            &'a self,
1033            _request: ModelRequest<'a>,
1034            _events: ModelEventSink,
1035        ) -> crate::BoxFuture<'a, crate::Result<ModelOutput>> {
1036            Box::pin(async { Err(crate::Error::Provider("unused".into())) })
1037        }
1038    }
1039
1040    #[test]
1041    fn request_input_is_borrowed_until_replaced() {
1042        let original = vec![Value::String("original".into())];
1043        let role = AgentRole::Main;
1044        let router = ModelRouter::new("test", Arc::new(NoModel));
1045        let mut context = ModelRequestContext {
1046            author: &MessageAuthor::User,
1047            role: &role,
1048            model: &router,
1049            provider: "test",
1050            session_id: "session",
1051            turn_id: "turn",
1052            model_step: 0,
1053            input: Cow::Borrowed(&original),
1054        };
1055
1056        assert!(matches!(&context.input, Cow::Borrowed(_)));
1057        context.replace_input(vec![Value::String("replacement".into())]);
1058        assert!(matches!(&context.input, Cow::Owned(_)));
1059        assert_eq!(original, [Value::String("original".into())]);
1060    }
1061}