Skip to main content

theway_core/agent/runtime_extensions/
mod.rs

1//! Engine-independent runtime-extension ports owned by core lifecycle seams.
2//!
3//! The embedding runtime implements these ports and translates an invocation to
4//! its extension engine. Core validates every returned ABI action batch before
5//! exposing a class-specific result to the lifecycle call site.
6
7mod compaction;
8mod context;
9mod message;
10mod request;
11mod run;
12mod scope;
13mod session;
14mod state;
15mod tool;
16
17use async_trait::async_trait;
18use serde_json::Value;
19use theway_contract::extension::{
20    ExtensionAction, ExtensionActionBatch, ExtensionErrorCode, ExtensionErrorEnvelope,
21    ExtensionGateDecision, ExtensionHookClass, ExtensionHookContract, ExtensionLifecycleEvent,
22    ExtensionModelRef, ExtensionScopeIds,
23};
24
25pub use compaction::RuntimeCompactionExtensionPort;
26pub use context::{
27    ExtensionModelContextItem, ExtensionModelContextProjection,
28    ExtensionModelContextProjectionError,
29};
30pub use message::RuntimeMessageExtensionPort;
31pub use request::RuntimeRequestExtensionPort;
32pub use run::RuntimeRunExtensionPort;
33pub use scope::{RuntimeExtensionScopeAllocator, RuntimeExtensionScopeKind, ScopeAllocationError};
34pub use session::RuntimeSessionExtensionPort;
35pub use state::{
36    NoopSessionExtensionStatePort, PersistentSessionExtensionStatePort, SessionExtensionStateError,
37    SessionExtensionStatePort,
38};
39pub use tool::RuntimeToolExtensionPort;
40
41pub type RawRuntimeExtensionResult = Result<ExtensionActionBatch, ExtensionErrorEnvelope>;
42pub type RuntimeExtensionResult = Result<ValidatedRuntimeExtensionResult, ExtensionErrorEnvelope>;
43
44#[derive(Clone, Debug, PartialEq, Eq)]
45pub struct RuntimeExtensionContext {
46    pub session_id: String,
47    pub cwd: String,
48    pub sequence: u64,
49    pub scope: ExtensionScopeIds,
50    pub model: Option<ExtensionModelRef>,
51    pub has_interactive_client: bool,
52    pub cancelled: bool,
53    pub deadline_unix_ms: Option<u64>,
54}
55
56impl RuntimeExtensionContext {
57    pub fn new(session_id: impl Into<String>, cwd: impl Into<String>, sequence: u64) -> Self {
58        Self {
59            session_id: session_id.into(),
60            cwd: cwd.into(),
61            sequence,
62            scope: ExtensionScopeIds::default(),
63            model: None,
64            has_interactive_client: false,
65            cancelled: false,
66            deadline_unix_ms: None,
67        }
68    }
69}
70
71#[derive(Clone, Debug, PartialEq)]
72pub struct RuntimeExtensionInvocation {
73    event: ExtensionLifecycleEvent,
74    class: ExtensionHookClass,
75    context: RuntimeExtensionContext,
76    payload: Value,
77}
78
79impl RuntimeExtensionInvocation {
80    pub fn new(
81        event: ExtensionLifecycleEvent,
82        class: ExtensionHookClass,
83        context: RuntimeExtensionContext,
84        payload: Value,
85    ) -> Result<Self, ExtensionErrorEnvelope> {
86        ExtensionHookContract::for_hook(event, class)?;
87        if context.session_id.trim().is_empty() || context.cwd.trim().is_empty() {
88            return Err(ExtensionErrorEnvelope::new(
89                ExtensionErrorCode::InvalidPayload,
90                "runtime extension context requires session_id and cwd",
91            ));
92        }
93        if context.sequence == 0 {
94            return Err(ExtensionErrorEnvelope::new(
95                ExtensionErrorCode::InvalidPayload,
96                "runtime extension lifecycle sequence must be greater than zero",
97            ));
98        }
99        if !payload.is_object() {
100            return Err(ExtensionErrorEnvelope::new(
101                ExtensionErrorCode::InvalidPayload,
102                "runtime extension event payload must be an object",
103            ));
104        }
105        Ok(Self {
106            event,
107            class,
108            context,
109            payload,
110        })
111    }
112
113    pub const fn event(&self) -> ExtensionLifecycleEvent {
114        self.event
115    }
116
117    pub const fn class(&self) -> ExtensionHookClass {
118        self.class
119    }
120
121    pub fn context(&self) -> &RuntimeExtensionContext {
122        &self.context
123    }
124
125    pub fn payload(&self) -> &Value {
126        &self.payload
127    }
128}
129
130#[derive(Clone, Debug, PartialEq, Eq)]
131pub struct ValidatedObserveResult;
132
133#[derive(Clone, Debug, PartialEq, Eq)]
134pub struct ValidatedTransformResult {
135    event: ExtensionLifecycleEvent,
136    actions: Vec<ExtensionAction>,
137}
138
139impl ValidatedTransformResult {
140    pub const fn event(&self) -> ExtensionLifecycleEvent {
141        self.event
142    }
143
144    pub fn actions(&self) -> &[ExtensionAction] {
145        &self.actions
146    }
147}
148
149#[derive(Clone, Debug, PartialEq, Eq)]
150pub struct ValidatedGateResult {
151    event: ExtensionLifecycleEvent,
152    decision: ExtensionGateDecision,
153    actions: Vec<ExtensionAction>,
154}
155
156impl ValidatedGateResult {
157    pub const fn event(&self) -> ExtensionLifecycleEvent {
158        self.event
159    }
160
161    pub fn decision(&self) -> &ExtensionGateDecision {
162        &self.decision
163    }
164
165    pub fn actions(&self) -> &[ExtensionAction] {
166        &self.actions
167    }
168}
169
170#[derive(Clone, Debug, PartialEq, Eq)]
171pub struct ValidatedRegisterResult {
172    actions: Vec<ExtensionAction>,
173}
174
175impl ValidatedRegisterResult {
176    pub fn actions(&self) -> &[ExtensionAction] {
177        &self.actions
178    }
179}
180
181#[derive(Clone, Debug, PartialEq, Eq)]
182pub enum ValidatedRuntimeExtensionResult {
183    Observe(ValidatedObserveResult),
184    Transform(ValidatedTransformResult),
185    Gate(ValidatedGateResult),
186    Register(ValidatedRegisterResult),
187}
188
189pub trait RuntimeExtensionPort:
190    RuntimeSessionExtensionPort
191    + RuntimeRunExtensionPort
192    + RuntimeRequestExtensionPort
193    + RuntimeMessageExtensionPort
194    + RuntimeToolExtensionPort
195    + RuntimeCompactionExtensionPort
196    + Send
197    + Sync
198{
199}
200
201impl<T> RuntimeExtensionPort for T where
202    T: RuntimeSessionExtensionPort
203        + RuntimeRunExtensionPort
204        + RuntimeRequestExtensionPort
205        + RuntimeMessageExtensionPort
206        + RuntimeToolExtensionPort
207        + RuntimeCompactionExtensionPort
208        + Send
209        + Sync
210{
211}
212
213#[derive(Clone, Copy, Debug, PartialEq, Eq)]
214pub enum RuntimeExtensionDomain {
215    Session,
216    Run,
217    Request,
218    Message,
219    Tool,
220    Compaction,
221}
222
223fn validate_domain_event(
224    domain: RuntimeExtensionDomain,
225    invocation: &RuntimeExtensionInvocation,
226) -> Result<(), ExtensionErrorEnvelope> {
227    let event = invocation.event();
228    let valid = match domain {
229        RuntimeExtensionDomain::Session => matches!(
230            event,
231            ExtensionLifecycleEvent::SessionStart
232                | ExtensionLifecycleEvent::BeforeSessionSwitch
233                | ExtensionLifecycleEvent::SessionSwitched
234                | ExtensionLifecycleEvent::BeforeSessionFork
235                | ExtensionLifecycleEvent::SessionForked
236                | ExtensionLifecycleEvent::SessionShutdown
237        ),
238        RuntimeExtensionDomain::Run => matches!(
239            event,
240            ExtensionLifecycleEvent::BeforeRun
241                | ExtensionLifecycleEvent::RunStarted
242                | ExtensionLifecycleEvent::TurnStarted
243                | ExtensionLifecycleEvent::TurnCompleted
244                | ExtensionLifecycleEvent::RunEnded
245                | ExtensionLifecycleEvent::RunError
246                | ExtensionLifecycleEvent::RunSettled
247        ),
248        RuntimeExtensionDomain::Request => matches!(
249            event,
250            ExtensionLifecycleEvent::Input
251                | ExtensionLifecycleEvent::BeforeModelSelection
252                | ExtensionLifecycleEvent::ModelSelected
253                | ExtensionLifecycleEvent::Context
254                | ExtensionLifecycleEvent::BeforeModelRequest
255                | ExtensionLifecycleEvent::BeforeProviderRequestHeaders
256                | ExtensionLifecycleEvent::BeforeProviderRequestRaw
257                | ExtensionLifecycleEvent::ProviderResponse
258                | ExtensionLifecycleEvent::ProviderRequestFailed
259        ),
260        RuntimeExtensionDomain::Message => matches!(
261            event,
262            ExtensionLifecycleEvent::MessageStart
263                | ExtensionLifecycleEvent::MessageUpdate
264                | ExtensionLifecycleEvent::MessageEnd
265        ),
266        RuntimeExtensionDomain::Tool => matches!(
267            event,
268            ExtensionLifecycleEvent::ToolCall
269                | ExtensionLifecycleEvent::ToolExecutionStart
270                | ExtensionLifecycleEvent::ToolExecutionUpdate
271                | ExtensionLifecycleEvent::ToolExecutionEnd
272                | ExtensionLifecycleEvent::ToolResult
273        ),
274        RuntimeExtensionDomain::Compaction => matches!(
275            event,
276            ExtensionLifecycleEvent::BeforeCompaction
277                | ExtensionLifecycleEvent::CompactionSucceeded
278                | ExtensionLifecycleEvent::CompactionFailed
279        ),
280    };
281    if valid {
282        Ok(())
283    } else {
284        Err(ExtensionErrorEnvelope::new(
285            ExtensionErrorCode::InvalidHook,
286            format!("event {event:?} does not belong to the {domain:?} core extension port"),
287        ))
288    }
289}
290
291fn validate_hook_result(
292    invocation: &RuntimeExtensionInvocation,
293    result: ExtensionActionBatch,
294) -> RuntimeExtensionResult {
295    let contract = ExtensionHookContract::for_hook(invocation.event, invocation.class)?;
296    contract.validate_result(&result)?;
297    Ok(match invocation.class {
298        ExtensionHookClass::Observe => {
299            ValidatedRuntimeExtensionResult::Observe(ValidatedObserveResult)
300        }
301        ExtensionHookClass::Transform => {
302            ValidatedRuntimeExtensionResult::Transform(ValidatedTransformResult {
303                event: invocation.event,
304                actions: result.actions,
305            })
306        }
307        ExtensionHookClass::Gate => ValidatedRuntimeExtensionResult::Gate(ValidatedGateResult {
308            event: invocation.event,
309            decision: result.decision.unwrap_or(ExtensionGateDecision::Abstain),
310            actions: result.actions,
311        }),
312        ExtensionHookClass::Register => {
313            ValidatedRuntimeExtensionResult::Register(ValidatedRegisterResult {
314                actions: result.actions,
315            })
316        }
317    })
318}
319
320#[derive(Clone, Copy, Debug, Default)]
321pub struct NoopRuntimeExtensionPort;
322
323fn empty_action_batch() -> ExtensionActionBatch {
324    ExtensionActionBatch {
325        decision: None,
326        actions: Vec::new(),
327    }
328}
329
330macro_rules! impl_noop_domain {
331    ($trait_name:ident, $method:ident) => {
332        #[async_trait]
333        impl $trait_name for NoopRuntimeExtensionPort {
334            async fn $method(
335                &self,
336                _invocation: RuntimeExtensionInvocation,
337            ) -> RawRuntimeExtensionResult {
338                Ok(empty_action_batch())
339            }
340        }
341    };
342}
343
344impl_noop_domain!(RuntimeSessionExtensionPort, invoke_session);
345impl_noop_domain!(RuntimeRunExtensionPort, invoke_run);
346impl_noop_domain!(RuntimeMessageExtensionPort, invoke_message);
347impl_noop_domain!(RuntimeToolExtensionPort, invoke_tool);
348impl_noop_domain!(RuntimeCompactionExtensionPort, invoke_compaction);
349
350#[async_trait]
351impl RuntimeRequestExtensionPort for NoopRuntimeExtensionPort {
352    fn has_request_hook(
353        &self,
354        _event: ExtensionLifecycleEvent,
355        _class: ExtensionHookClass,
356    ) -> bool {
357        false
358    }
359
360    async fn invoke_request(
361        &self,
362        _invocation: RuntimeExtensionInvocation,
363    ) -> RawRuntimeExtensionResult {
364        Ok(empty_action_batch())
365    }
366}
367
368#[cfg(test)]
369tests_bridge_macro::tests_bridge!("agent/runtime_extensions");