theway_core/agent/runtime_extensions/
mod.rs1mod 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");