use std::collections::BTreeMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use serde_json::Value;
use theway_contract::extension::{
ExtensionActionBatch, ExtensionCatalogEntry, ExtensionCatalogStatus, ExtensionDiagnostic,
ExtensionDiagnosticCode, ExtensionHookClass, ExtensionHookContract, ExtensionLifecycleEvent,
};
use theway_core::AgentTool;
use theway_core::agent::runtime_extensions::{
ExtensionModelContextProjection, RawRuntimeExtensionResult, RuntimeExtensionInvocation,
};
use super::catalog::{ExtensionPackage, PackageCatalog};
use super::compaction::LegacyCompactionHost;
use super::diagnostics;
use super::dispatch_result::{
accept_transform_batch, decode_batch, empty_batch, failed_gate_decision, merge_batch,
validate_ephemeral_actions,
};
use super::dispatcher::{self, HookRegistration, RuntimeExtensionHostConfig};
use super::effects::{InstanceHealth, InstanceLifecyclePhase};
use super::engine::{EngineInstanceKey, QuickJsEnginePool};
use super::observation::{ObservationDispatch, ObservationJob, ObservationQueue, diagnostic_code};
use super::registration_runtime::RegistrationRuntime;
use super::reload::HostReloadState;
use super::state::HostLifecycleSequence;
use super::state_runtime::ExtensionStateRuntime;
#[derive(Clone, Debug, PartialEq)]
pub struct ExtensionInvocationOutput {
pub extension_id: String,
pub value: Value,
}
#[derive(Clone)]
pub(super) struct ActiveExtension {
pub(super) package: Arc<ExtensionPackage>,
pub(super) key: EngineInstanceKey,
pub(super) registrations: Vec<HookRegistration>,
pub(super) phase: InstanceLifecyclePhase,
pub(super) health: Arc<InstanceHealth>,
pub(super) observation_queues: BTreeMap<u64, Arc<ObservationQueue>>,
}
pub struct SessionPluginHost {
pub(super) session_id: String,
pub(super) cwd: String,
pub(super) engine: QuickJsEnginePool,
pub(super) config: RuntimeExtensionHostConfig,
pub(super) sequence: Arc<HostLifecycleSequence>,
pub(super) active: tokio::sync::Mutex<Vec<ActiveExtension>>,
pub(super) catalog: Arc<parking_lot::RwLock<PackageCatalog>>,
pub(super) diagnostics: Arc<parking_lot::Mutex<Vec<ExtensionDiagnostic>>>,
pub(super) subscription_counts:
parking_lot::RwLock<BTreeMap<(ExtensionLifecycleEvent, ExtensionHookClass), usize>>,
pub(super) shutdown: Arc<AtomicBool>,
pub(super) registration_runtime: RegistrationRuntime,
pub(super) state_runtime: ExtensionStateRuntime,
pub(super) reload_state: HostReloadState,
pub(super) legacy_compaction: Option<Arc<LegacyCompactionHost>>,
pub(super) reload_catalog: Option<Arc<parking_lot::RwLock<PackageCatalog>>>,
pub(super) reload_base_tools: parking_lot::Mutex<Vec<Arc<dyn AgentTool>>>,
pub(super) reload_tool_publisher: parking_lot::Mutex<Option<ReloadToolPublisher>>,
}
pub(super) type ReloadToolPublisher = Arc<dyn Fn(Vec<Arc<dyn AgentTool>>) + Send + Sync>;
impl SessionPluginHost {
pub async fn invoke(
&self,
event: ExtensionLifecycleEvent,
payload: Value,
) -> Vec<ExtensionInvocationOutput> {
if self.shutdown.load(Ordering::Acquire) {
return Vec::new();
}
let active = self.active.lock().await;
let mut outputs = Vec::with_capacity(active.len());
let mut index = 0;
while index < active.len() {
if active[index].health.is_open() {
index += 1;
continue;
}
if !active[index]
.registrations
.iter()
.any(|registration| registration.event == event)
{
index += 1;
continue;
}
match self
.invoke_extension(&active[index], event, payload.clone())
.await
{
Ok(value) => {
active[index].health.record_success();
outputs.push(ExtensionInvocationOutput {
extension_id: active[index].package.manifest().id.clone(),
value,
});
index += 1;
}
Err(error) => {
self.record_hook_failure(
&active[index],
event,
ExtensionDiagnosticCode::HookFailed,
error,
)
.await;
index += 1;
}
}
}
outputs
}
pub async fn shutdown(&self) {
if self.shutdown.swap(true, Ordering::AcqRel) {
return;
}
let mut active = self.active.lock().await;
for mut extension in active.drain(..) {
if extension.phase == InstanceLifecyclePhase::Started {
self.invoke_cleanup_event(
&extension,
ExtensionLifecycleEvent::SessionShutdown,
serde_json::json!({"reason": "shutdown"}),
)
.await;
}
self.invoke_cleanup_event(
&extension,
ExtensionLifecycleEvent::ExtensionUnload,
serde_json::json!({"reason": "shutdown"}),
)
.await;
self.dispose_extension_effects(&extension.package.manifest().id);
self.engine.dispose(&extension.key).await;
self.remove_subscriptions(&extension.registrations);
extension.phase = InstanceLifecyclePhase::Disposed;
}
}
pub async fn active_extension_ids(&self) -> Vec<String> {
self.active
.lock()
.await
.iter()
.filter(|extension| !extension.health.is_open())
.map(|extension| extension.package.manifest().id.clone())
.collect()
}
pub async fn active_effect_count(&self) -> usize {
self.registration_runtime.active_count()
}
pub fn model_context_projection(&self) -> ExtensionModelContextProjection {
self.state_runtime.model_context()
}
pub fn catalog_entries(&self) -> Vec<ExtensionCatalogEntry> {
self.catalog.read().entries().to_vec()
}
pub fn diagnostics(&self) -> Vec<ExtensionDiagnostic> {
let mut diagnostics = self.diagnostics.lock().clone();
diagnostics.extend(self.engine.broker_diagnostics(&self.session_id));
diagnostics
}
pub fn audit_events(&self) -> Vec<theway_contract::extension::ExtensionAuditEvent> {
self.engine
.audit_log()
.events()
.into_iter()
.filter(|event| event.session_id.as_deref() == Some(self.session_id.as_str()))
.collect()
}
pub(super) async fn invoke_extension(
&self,
extension: &ActiveExtension,
event: ExtensionLifecycleEvent,
payload: Value,
) -> Result<Value, String> {
let mut aggregate = empty_batch();
for registration in extension
.registrations
.iter()
.filter(|registration| registration.event == event)
{
if !self.registration_runtime.is_registration_active(
&self.effect_owner(&extension.key.extension_id),
registration.registration_id,
) {
continue;
}
if !registration.accepts_payload(&payload) {
return Err("extension event payload does not match the hook payloadSchema".into());
}
let (result, origin_sequence) = self
.invoke_registration(&extension.key, registration, event, payload.clone())
.await?;
let mut batch = decode_batch(result.value)?;
batch.actions.extend(result.queued_durable_actions);
if batch.actions.len() > self.config.max_actions {
return Err("extension action count exceeds the configured limit".into());
}
registration
.contract
.validate_result(&batch)
.map_err(|error| error.message)?;
dispatcher::validate_action_capabilities(
&batch,
extension.package.granted_permissions(),
)?;
validate_ephemeral_actions(event, registration.class, &payload, &batch)?;
let emitted = self.emitted_diagnostics(
&extension.package.manifest().id,
event,
origin_sequence,
&batch,
)?;
self.state_runtime
.commit_batch(
&extension.package.manifest().id,
origin_sequence,
&mut batch,
)
.await?;
self.diagnostics.lock().extend(emitted);
if merge_batch(event, registration.class, &mut aggregate, batch) {
break;
}
}
serde_json::to_value(aggregate)
.map_err(|error| format!("extension action batch serialization failed: {error}"))
}
fn emitted_diagnostics(
&self,
extension_id: &str,
event: ExtensionLifecycleEvent,
sequence: u64,
batch: &theway_contract::extension::ExtensionActionBatch,
) -> Result<Vec<ExtensionDiagnostic>, String> {
batch
.actions
.iter()
.filter(|action| {
action.kind == theway_contract::extension::ExtensionActionKind::EmitDiagnostic
})
.map(|action| {
diagnostics::emitted(
extension_id,
&self.session_id,
event,
sequence,
action.payload.clone(),
)
})
.collect()
}
pub(super) async fn cleanup_failed_start(
&self,
extension: &ActiveExtension,
session_started: bool,
) {
if session_started {
self.invoke_cleanup_event(
extension,
ExtensionLifecycleEvent::SessionShutdown,
serde_json::json!({"reason": "initialization_failed"}),
)
.await;
}
self.invoke_cleanup_event(
extension,
ExtensionLifecycleEvent::ExtensionUnload,
serde_json::json!({"reason": "initialization_failed"}),
)
.await;
self.dispose_extension_effects(&extension.package.manifest().id);
self.engine.dispose(&extension.key).await;
}
async fn record_hook_failure(
&self,
extension: &ActiveExtension,
event: ExtensionLifecycleEvent,
code: ExtensionDiagnosticCode,
error: String,
) {
self.diagnostics.lock().push(diagnostics::invocation(
extension.package.manifest().id.clone(),
self.session_id.clone(),
event,
code,
format!("extension hook failed: {error}"),
));
if code != ExtensionDiagnosticCode::Cancelled
&& extension
.health
.record_failure(self.config.circuit_failure_threshold)
{
self.catalog.write().set_effective_status(
&extension.package.manifest().id,
ExtensionCatalogStatus::Disabled,
Some(ExtensionDiagnosticCode::CircuitOpened),
);
self.diagnostics.lock().push(diagnostics::circuit_opened(
extension.package.manifest().id.clone(),
self.session_id.clone(),
));
self.dispose_extension_effects(&extension.package.manifest().id);
self.engine.dispose(&extension.key).await;
}
}
pub(super) fn record_load_fault(&self, package: &ExtensionPackage, error: String) {
self.diagnostics.lock().push(diagnostics::faulted(
package.manifest().id.clone(),
self.session_id.clone(),
format!("extension load failed: {error}"),
));
self.catalog.write().set_effective_status(
&package.manifest().id,
ExtensionCatalogStatus::Faulted,
Some(ExtensionDiagnosticCode::LoadFailed),
);
}
pub(super) fn record_state_migration_fault(&self, package: &ExtensionPackage, error: String) {
self.diagnostics.lock().push(diagnostics::invocation(
package.manifest().id.clone(),
self.session_id.clone(),
ExtensionLifecycleEvent::ExtensionLoad,
ExtensionDiagnosticCode::StateMigrationFailed,
format!("extension state reconstruction or migration failed: {error}"),
));
self.catalog.write().set_effective_status(
&package.manifest().id,
ExtensionCatalogStatus::Disabled,
Some(ExtensionDiagnosticCode::StateMigrationFailed),
);
}
pub(super) async fn invoke_cleanup_event(
&self,
extension: &ActiveExtension,
event: ExtensionLifecycleEvent,
payload: Value,
) {
if !extension
.registrations
.iter()
.any(|registration| registration.event == event)
{
return;
}
if let Err(error) = self.invoke_extension(extension, event, payload).await {
self.diagnostics.lock().push(diagnostics::hook_failed(
extension.package.manifest().id.clone(),
self.session_id.clone(),
format!("extension cleanup hook failed: {error}"),
));
}
}
pub(super) fn add_subscriptions(&self, registrations: &[HookRegistration]) {
let mut counts = self.subscription_counts.write();
for registration in registrations {
*counts
.entry((registration.event, registration.class))
.or_default() += 1;
}
}
pub(super) fn remove_subscriptions(&self, registrations: &[HookRegistration]) {
let mut counts = self.subscription_counts.write();
for registration in registrations {
let key = (registration.event, registration.class);
if let Some(count) = counts.get_mut(&key) {
*count -= 1;
if *count == 0 {
counts.remove(&key);
}
}
}
}
pub(super) fn has_subscription(
&self,
event: ExtensionLifecycleEvent,
class: ExtensionHookClass,
) -> bool {
self.subscription_counts
.read()
.contains_key(&(event, class))
}
pub(super) async fn invoke_runtime(
&self,
invocation: RuntimeExtensionInvocation,
) -> RawRuntimeExtensionResult {
if self.shutdown.load(Ordering::Acquire) {
return Ok(empty_batch());
}
let event = invocation.event();
let class = invocation.class();
ExtensionHookContract::for_hook(event, class)?;
if !self.has_subscription(event, class) && !self.has_request_registration(event, class) {
return Ok(empty_batch());
}
let mut aggregate = empty_batch();
let mut current_payload = invocation.payload().clone();
if event == ExtensionLifecycleEvent::BeforeModelRequest
&& class == ExtensionHookClass::Transform
{
self.apply_request_registrations(&invocation, &mut current_payload, &mut aggregate)
.await;
}
let active = self.active.lock().await.clone();
for extension in &active {
let registrations: Vec<_> = extension
.registrations
.iter()
.filter(|registration| registration.event == event && registration.class == class)
.cloned()
.collect();
if registrations.is_empty() {
continue;
}
if extension.health.is_open() {
if class == ExtensionHookClass::Gate {
aggregate.decision = Some(failed_gate_decision());
break;
}
continue;
}
for registration in registrations {
if !self.registration_runtime.is_registration_active(
&self.effect_owner(&extension.package.manifest().id),
registration.registration_id,
) {
continue;
}
if registration.delivery
== theway_contract::extension::ExtensionDeliveryPolicy::BoundedCoalescing
{
self.enqueue_observation(extension, registration, &invocation);
continue;
}
let result = self
.dispatch_registration(
extension,
®istration,
&invocation,
current_payload.clone(),
)
.await;
match result {
Ok(batch) => {
let accepted = if class == ExtensionHookClass::Transform {
accept_transform_batch(
event,
&mut current_payload,
&mut aggregate,
batch,
)
} else {
Ok(merge_batch(event, class, &mut aggregate, batch))
};
match accepted {
Ok(stop) => {
extension.health.record_success();
if stop {
return Ok(aggregate);
}
}
Err(error) => {
self.record_hook_failure(
extension,
event,
ExtensionDiagnosticCode::ContractViolation,
error,
)
.await;
if registration.failure
== theway_contract::extension::ExtensionHookFailurePolicy::Deny
{
aggregate.decision = Some(failed_gate_decision());
return Ok(aggregate);
}
}
}
}
Err((code, error)) => {
self.record_hook_failure(extension, event, code, error)
.await;
if registration.failure
== theway_contract::extension::ExtensionHookFailurePolicy::Deny
{
aggregate.decision = Some(failed_gate_decision());
return Ok(aggregate);
}
}
}
}
}
if event == ExtensionLifecycleEvent::SessionStart {
let mut active = self.active.lock().await;
for extension in active.iter_mut() {
extension.phase = InstanceLifecyclePhase::Started;
}
}
Ok(aggregate)
}
async fn dispatch_registration(
&self,
extension: &ActiveExtension,
registration: &HookRegistration,
invocation: &RuntimeExtensionInvocation,
payload: Value,
) -> Result<ExtensionActionBatch, (ExtensionDiagnosticCode, String)> {
if !registration.accepts_payload(&payload) {
return Err((
ExtensionDiagnosticCode::ContractViolation,
"extension event payload does not match the hook payloadSchema".into(),
));
}
if invocation.context().cancelled || self.shutdown.load(Ordering::Acquire) {
return Err((
ExtensionDiagnosticCode::Cancelled,
"extension invocation was cancelled".into(),
));
}
let envelope = dispatcher::runtime_envelope_with_payload(
&extension.package.manifest().id,
invocation,
payload,
);
let result = self
.engine
.invoke_controlled_with_effects(
&extension.key,
&envelope,
registration.registration_id,
self.config.deadline(registration.deadline),
Arc::clone(&self.shutdown),
self.config.broker_operation_quota,
)
.await
.map_err(|error| (diagnostic_code(error.kind), error.message))?;
self.registration_runtime.apply_disposals(
&self.effect_owner(&extension.package.manifest().id),
&result.disposed_registration_ids,
);
if invocation.context().cancelled || self.shutdown.load(Ordering::Acquire) {
return Err((
ExtensionDiagnosticCode::Cancelled,
"extension result arrived after cancellation".into(),
));
}
let mut batch = decode_batch(result.value)
.map_err(|error| (ExtensionDiagnosticCode::ContractViolation, error))?;
batch.actions.extend(result.queued_durable_actions);
if batch.actions.len() > self.config.max_actions {
return Err((
ExtensionDiagnosticCode::ResourceLimit,
"extension action count exceeds the configured limit".into(),
));
}
registration
.contract
.validate_result(&batch)
.map_err(|error| (ExtensionDiagnosticCode::ContractViolation, error.message))?;
dispatcher::validate_action_capabilities(&batch, extension.package.granted_permissions())
.map_err(|error| (ExtensionDiagnosticCode::PermissionDenied, error))?;
validate_ephemeral_actions(
invocation.event(),
registration.class,
&envelope.payload,
&batch,
)
.map_err(|error| (ExtensionDiagnosticCode::ContractViolation, error))?;
let emitted = self
.emitted_diagnostics(
&extension.package.manifest().id,
invocation.event(),
invocation.context().sequence,
&batch,
)
.map_err(|error| (ExtensionDiagnosticCode::ContractViolation, error))?;
self.state_runtime
.commit_batch(
&extension.package.manifest().id,
invocation.context().sequence,
&mut batch,
)
.await
.map_err(|error| (ExtensionDiagnosticCode::HookFailed, error))?;
self.diagnostics.lock().extend(emitted);
Ok(batch)
}
fn enqueue_observation(
&self,
extension: &ActiveExtension,
registration: HookRegistration,
invocation: &RuntimeExtensionInvocation,
) {
if !registration.accepts_payload(invocation.payload()) {
self.diagnostics.lock().push(diagnostics::invocation(
extension.package.manifest().id.clone(),
self.session_id.clone(),
registration.event,
ExtensionDiagnosticCode::ContractViolation,
"extension event payload does not match the hook payloadSchema",
));
return;
}
debug_assert_eq!(
registration.failure,
theway_contract::extension::ExtensionHookFailurePolicy::Continue
);
let Some(queue) = extension
.observation_queues
.get(®istration.registration_id)
.cloned()
else {
return;
};
let job = ObservationJob {
envelope: dispatcher::runtime_envelope(&extension.package.manifest().id, invocation),
cancellation: Arc::clone(&self.shutdown),
};
let Some(first) = queue.enqueue(job) else {
return;
};
ObservationDispatch {
extension_id: extension.package.manifest().id.clone(),
session_id: self.session_id.clone(),
key: extension.key.clone(),
registration,
engine: self.engine.clone(),
config: self.config.clone(),
health: Arc::clone(&extension.health),
diagnostics: Arc::clone(&self.diagnostics),
catalog: Arc::clone(&self.catalog),
registration_runtime: self.registration_runtime.clone(),
}
.spawn(queue, first);
}
pub(super) async fn unload_after_core_shutdown(&self) {
if self.shutdown.swap(true, Ordering::AcqRel) {
return;
}
let mut active = self.active.lock().await;
for mut extension in active.drain(..) {
self.invoke_cleanup_event(
&extension,
ExtensionLifecycleEvent::ExtensionUnload,
serde_json::json!({"reason": "shutdown"}),
)
.await;
self.dispose_extension_effects(&extension.package.manifest().id);
self.engine.dispose(&extension.key).await;
self.remove_subscriptions(&extension.registrations);
extension.phase = InstanceLifecyclePhase::Disposed;
}
}
}