1use std::collections::BTreeMap;
2use std::sync::{Arc, RwLock};
3
4use sha2::{Digest, Sha256};
5
6use crate::approval::{ApprovalBroker, ApprovalProvider};
7use crate::budget::HostCostMeter;
8use crate::checkpoint::{CheckpointExtension, IdempotentRunEventStore, ReconciliationProvider};
9use crate::llm::LlmClient;
10use crate::memory::MemoryProvider;
11use crate::runtime::engine::RunEventHandler;
12use crate::runtime::hooks::RuntimeHook;
13use crate::runtime::lifecycle::AfterCycleHook;
14use crate::runtime::state::CheckpointStore;
15use crate::runtime::sub_task_manager::SubTaskManager;
16use crate::runtime::CancellationToken;
17use crate::tools::{ApprovalPolicy, CanUseToolPredicate, ToolPolicy, ToolRegistry};
18use crate::workspace::WorkspaceBackend;
19
20use super::{
21 CapabilityRef, DistributedCapabilities, DistributedCheckpointExtensionRef,
22 DistributedToolPolicy, ToolsetRef,
23};
24
25#[derive(Debug, Clone, PartialEq, Eq)]
26pub struct DistributedCapabilityError {
27 message: String,
28}
29
30impl DistributedCapabilityError {
31 fn new(message: impl Into<String>) -> Self {
32 Self {
33 message: message.into(),
34 }
35 }
36}
37
38impl std::fmt::Display for DistributedCapabilityError {
39 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 formatter.write_str(&self.message)
41 }
42}
43
44impl std::error::Error for DistributedCapabilityError {}
45
46type CapabilityKey = (String, String);
47
48#[derive(Default)]
49struct CapabilityMaps {
50 toolsets: BTreeMap<CapabilityKey, ToolRegistry>,
51 llm_clients: BTreeMap<CapabilityKey, Arc<dyn LlmClient>>,
52 workspace_backends: BTreeMap<CapabilityKey, Arc<dyn WorkspaceBackend>>,
53 approval_providers: BTreeMap<CapabilityKey, Arc<dyn ApprovalProvider>>,
54 approval_brokers: BTreeMap<CapabilityKey, ApprovalBroker>,
55 cancellations: BTreeMap<CapabilityKey, CancellationToken>,
56 event_sinks: BTreeMap<CapabilityKey, RunEventHandler>,
57 host_cost_meters: BTreeMap<CapabilityKey, Arc<dyn HostCostMeter>>,
58 app_states: BTreeMap<CapabilityKey, Arc<dyn std::any::Any + Send + Sync>>,
59 memory_providers: BTreeMap<CapabilityKey, Arc<dyn MemoryProvider>>,
60 hooks: BTreeMap<CapabilityKey, Arc<dyn RuntimeHook>>,
61 after_cycle_hooks: BTreeMap<CapabilityKey, Arc<dyn AfterCycleHook>>,
62 observers: BTreeMap<CapabilityKey, RunEventHandler>,
63 sub_task_managers: BTreeMap<CapabilityKey, SubTaskManager>,
64 tool_predicates: BTreeMap<CapabilityKey, CanUseToolPredicate>,
65 checkpoint_stores: BTreeMap<CapabilityKey, Arc<dyn CheckpointStore>>,
66 checkpoint_event_stores: BTreeMap<CapabilityKey, Arc<dyn IdempotentRunEventStore>>,
67 checkpoint_extensions: BTreeMap<CapabilityKey, Arc<dyn CheckpointExtension>>,
68 reconciliation_providers: BTreeMap<CapabilityKey, Arc<dyn ReconciliationProvider>>,
69}
70
71#[derive(Clone)]
72pub struct DistributedCapabilityRegistry {
73 inner: Arc<RwLock<CapabilityMaps>>,
74}
75
76impl Default for DistributedCapabilityRegistry {
77 fn default() -> Self {
78 Self::new()
79 }
80}
81
82impl DistributedCapabilityRegistry {
83 pub fn new() -> Self {
84 let registry = Self::empty();
85 registry
86 .register_toolset(
87 ToolsetRef::default(),
88 crate::tools::build_default_registry(),
89 )
90 .expect("built-in tool schema digest is a compile-time parity contract");
91 registry
92 }
93
94 pub fn empty() -> Self {
95 Self {
96 inner: Arc::new(RwLock::new(CapabilityMaps::default())),
97 }
98 }
99
100 pub fn register_toolset(
101 &self,
102 reference: ToolsetRef,
103 registry: ToolRegistry,
104 ) -> Result<(), DistributedCapabilityError> {
105 let actual = toolset_schema_digest(®istry)?;
106 if actual != reference.schema_digest {
107 return Err(DistributedCapabilityError::new(format!(
108 "toolset {}@{} schema digest mismatch: expected {}, got {actual}",
109 reference.id, reference.version, reference.schema_digest
110 )));
111 }
112 self.write()?
113 .toolsets
114 .insert(key(&reference.capability_ref()), registry);
115 Ok(())
116 }
117
118 pub fn register_llm_client(&self, reference: CapabilityRef, client: Arc<dyn LlmClient>) {
119 self.write_unpoisoned()
120 .llm_clients
121 .insert(key(&reference), client);
122 }
123
124 pub fn register_workspace_backend(
125 &self,
126 reference: CapabilityRef,
127 backend: Arc<dyn WorkspaceBackend>,
128 ) {
129 self.write_unpoisoned()
130 .workspace_backends
131 .insert(key(&reference), backend);
132 }
133
134 pub fn register_approval_provider(
135 &self,
136 reference: CapabilityRef,
137 provider: Arc<dyn ApprovalProvider>,
138 ) {
139 self.write_unpoisoned()
140 .approval_providers
141 .insert(key(&reference), provider);
142 }
143
144 pub fn register_approval_broker(&self, reference: CapabilityRef, broker: ApprovalBroker) {
145 self.write_unpoisoned()
146 .approval_brokers
147 .insert(key(&reference), broker);
148 }
149
150 pub fn register_cancellation(&self, reference: CapabilityRef, token: CancellationToken) {
151 self.write_unpoisoned()
152 .cancellations
153 .insert(key(&reference), token);
154 }
155
156 pub fn register_event_sink(&self, reference: CapabilityRef, sink: RunEventHandler) {
157 self.write_unpoisoned()
158 .event_sinks
159 .insert(key(&reference), sink);
160 }
161
162 pub fn register_host_cost_meter(
163 &self,
164 reference: CapabilityRef,
165 meter: Arc<dyn HostCostMeter>,
166 ) {
167 self.write_unpoisoned()
168 .host_cost_meters
169 .insert(key(&reference), meter);
170 }
171
172 pub fn register_app_state(
173 &self,
174 reference: CapabilityRef,
175 state: Arc<dyn std::any::Any + Send + Sync>,
176 ) {
177 self.write_unpoisoned()
178 .app_states
179 .insert(key(&reference), state);
180 }
181
182 pub fn register_memory_provider(
183 &self,
184 reference: CapabilityRef,
185 provider: Arc<dyn MemoryProvider>,
186 ) {
187 self.write_unpoisoned()
188 .memory_providers
189 .insert(key(&reference), provider);
190 }
191
192 pub fn register_hook(&self, reference: CapabilityRef, hook: Arc<dyn RuntimeHook>) {
193 self.write_unpoisoned().hooks.insert(key(&reference), hook);
194 }
195
196 pub fn register_after_cycle_hook(
197 &self,
198 reference: CapabilityRef,
199 hook: Arc<dyn AfterCycleHook>,
200 ) {
201 self.write_unpoisoned()
202 .after_cycle_hooks
203 .insert(key(&reference), hook);
204 }
205
206 pub fn register_observer(&self, reference: CapabilityRef, observer: RunEventHandler) {
207 self.write_unpoisoned()
208 .observers
209 .insert(key(&reference), observer);
210 }
211
212 pub fn register_sub_task_manager(&self, reference: CapabilityRef, manager: SubTaskManager) {
213 self.write_unpoisoned()
214 .sub_task_managers
215 .insert(key(&reference), manager);
216 }
217
218 pub fn register_tool_predicate(
219 &self,
220 reference: CapabilityRef,
221 predicate: CanUseToolPredicate,
222 ) {
223 self.write_unpoisoned()
224 .tool_predicates
225 .insert(key(&reference), predicate);
226 }
227
228 pub fn register_checkpoint_store(
229 &self,
230 reference: CapabilityRef,
231 store: Arc<dyn CheckpointStore>,
232 ) {
233 self.write_unpoisoned()
234 .checkpoint_stores
235 .insert(key(&reference), store);
236 }
237
238 pub fn register_checkpoint_event_store(
239 &self,
240 reference: CapabilityRef,
241 store: Arc<dyn IdempotentRunEventStore>,
242 ) {
243 self.write_unpoisoned()
244 .checkpoint_event_stores
245 .insert(key(&reference), store);
246 }
247
248 pub fn register_checkpoint_extension(
249 &self,
250 reference: CapabilityRef,
251 extension: Arc<dyn CheckpointExtension>,
252 ) {
253 self.write_unpoisoned()
254 .checkpoint_extensions
255 .insert(key(&reference), extension);
256 }
257
258 pub fn register_reconciliation_provider(
259 &self,
260 reference: CapabilityRef,
261 provider: Arc<dyn ReconciliationProvider>,
262 ) {
263 self.write_unpoisoned()
264 .reconciliation_providers
265 .insert(key(&reference), provider);
266 }
267
268 pub(crate) fn resolve_checkpoint_store_required(
269 &self,
270 reference: &CapabilityRef,
271 ) -> Result<Arc<dyn CheckpointStore>, DistributedCapabilityError> {
272 self.read()?
273 .checkpoint_stores
274 .get(&key(reference))
275 .cloned()
276 .ok_or_else(|| unknown("checkpoint_store", reference))
277 }
278
279 pub fn resolve(
280 &self,
281 capabilities: &DistributedCapabilities,
282 ) -> Result<ResolvedDistributedCapabilities, DistributedCapabilityError> {
283 capabilities
284 .validate()
285 .map_err(DistributedCapabilityError::new)?;
286 let maps = self.read()?;
287 let tool_registry = resolve_toolset(&maps, &capabilities.toolset_ref)?;
288 let tool_policy = resolve_tool_policy(&maps, &capabilities.tool_policy)?;
289 let llm_client = optional(
290 &maps.llm_clients,
291 "llm_client",
292 &capabilities.llm_client_ref,
293 )?;
294 let workspace_backend = optional(
295 &maps.workspace_backends,
296 "workspace_backend",
297 &capabilities.workspace_backend_ref,
298 )?;
299 let approval_provider = optional(
300 &maps.approval_providers,
301 "approval_provider",
302 &capabilities.approval_provider_ref,
303 )?;
304 let approval_broker = optional(
305 &maps.approval_brokers,
306 "approval_broker",
307 &capabilities.approval_broker_ref,
308 )?;
309 let cancellation = optional(
310 &maps.cancellations,
311 "cancellation",
312 &capabilities.cancellation_ref,
313 )?;
314 let event_sink = optional(
315 &maps.event_sinks,
316 "event_sink",
317 &capabilities.event_sink_ref,
318 )?;
319 let host_cost_meter = optional(
320 &maps.host_cost_meters,
321 "host_cost_meter",
322 &capabilities.host_cost_meter_ref,
323 )?;
324 let app_state = optional(&maps.app_states, "app_state", &capabilities.app_state_ref)?;
325 let sub_task_manager = optional(
326 &maps.sub_task_managers,
327 "sub_task_manager",
328 &capabilities.sub_task_manager_ref,
329 )?;
330 let memory_providers = required_many(
331 &maps.memory_providers,
332 "memory_provider",
333 &capabilities.memory_provider_refs,
334 )?;
335 let hooks = required_many(&maps.hooks, "hook", &capabilities.hook_refs)?;
336 let after_cycle_hooks = required_many(
337 &maps.after_cycle_hooks,
338 "after_cycle_hook",
339 &capabilities.after_cycle_hook_refs,
340 )?;
341 let observers = required_many(&maps.observers, "observer", &capabilities.observer_refs)?;
342 let checkpoint_store = optional(
343 &maps.checkpoint_stores,
344 "checkpoint_store",
345 &capabilities.checkpoint_store_ref,
346 )?;
347 let checkpoint_event_store = optional(
348 &maps.checkpoint_event_stores,
349 "checkpoint_event_store",
350 &capabilities.checkpoint_event_store_ref,
351 )?;
352 let checkpoint_extensions = capabilities
353 .checkpoint_extension_refs
354 .iter()
355 .map(|descriptor| resolve_checkpoint_extension(&maps, descriptor))
356 .collect::<Result<Vec<_>, _>>()?;
357 let reconciliation_provider = optional(
358 &maps.reconciliation_providers,
359 "reconciliation_provider",
360 &capabilities.reconciliation_provider_ref,
361 )?;
362 Ok(ResolvedDistributedCapabilities {
363 tool_registry,
364 tool_policy,
365 llm_client,
366 workspace_backend,
367 approval_provider,
368 approval_broker,
369 approval_timeout_seconds: capabilities.approval_timeout_seconds,
370 cancellation,
371 event_sink,
372 host_cost_meter,
373 app_state,
374 memory_providers,
375 hooks,
376 after_cycle_hooks,
377 observers,
378 sub_task_manager,
379 checkpoint_store,
380 checkpoint_event_store,
381 checkpoint_extensions,
382 reconciliation_provider,
383 })
384 }
385
386 fn read(
387 &self,
388 ) -> Result<std::sync::RwLockReadGuard<'_, CapabilityMaps>, DistributedCapabilityError> {
389 self.inner.read().map_err(|_| {
390 DistributedCapabilityError::new("distributed capability registry lock poisoned")
391 })
392 }
393
394 fn write(
395 &self,
396 ) -> Result<std::sync::RwLockWriteGuard<'_, CapabilityMaps>, DistributedCapabilityError> {
397 self.inner.write().map_err(|_| {
398 DistributedCapabilityError::new("distributed capability registry lock poisoned")
399 })
400 }
401
402 fn write_unpoisoned(&self) -> std::sync::RwLockWriteGuard<'_, CapabilityMaps> {
403 self.inner
404 .write()
405 .unwrap_or_else(std::sync::PoisonError::into_inner)
406 }
407}
408
409#[derive(Clone)]
410pub struct ResolvedDistributedCapabilities {
411 pub tool_registry: ToolRegistry,
412 pub tool_policy: ToolPolicy,
413 pub llm_client: Option<Arc<dyn LlmClient>>,
414 pub workspace_backend: Option<Arc<dyn WorkspaceBackend>>,
415 pub approval_provider: Option<Arc<dyn ApprovalProvider>>,
416 pub approval_broker: Option<ApprovalBroker>,
417 pub approval_timeout_seconds: Option<f64>,
418 pub cancellation: Option<CancellationToken>,
419 pub event_sink: Option<RunEventHandler>,
420 pub host_cost_meter: Option<Arc<dyn HostCostMeter>>,
421 pub app_state: Option<Arc<dyn std::any::Any + Send + Sync>>,
422 pub memory_providers: Vec<Arc<dyn MemoryProvider>>,
423 pub hooks: Vec<Arc<dyn RuntimeHook>>,
424 pub after_cycle_hooks: Vec<Arc<dyn AfterCycleHook>>,
425 pub observers: Vec<RunEventHandler>,
426 pub sub_task_manager: Option<SubTaskManager>,
427 pub checkpoint_store: Option<Arc<dyn CheckpointStore>>,
428 pub checkpoint_event_store: Option<Arc<dyn IdempotentRunEventStore>>,
429 pub checkpoint_extensions: Vec<ResolvedDistributedCheckpointExtension>,
430 pub reconciliation_provider: Option<Arc<dyn ReconciliationProvider>>,
431}
432
433#[derive(Clone)]
434pub struct ResolvedDistributedCheckpointExtension {
435 pub descriptor: DistributedCheckpointExtensionRef,
436 pub extension: Arc<dyn CheckpointExtension>,
437}
438
439pub fn toolset_schema_digest(
440 registry: &ToolRegistry,
441) -> Result<String, DistributedCapabilityError> {
442 let schemas = registry
443 .list_openai_schemas(None)
444 .map_err(DistributedCapabilityError::new)?;
445 let canonical = serde_json::to_vec(&schemas).map_err(|error| {
446 DistributedCapabilityError::new(format!("failed to serialize toolset schemas: {error}"))
447 })?;
448 Ok(format!("{:x}", Sha256::digest(canonical)))
449}
450
451fn resolve_toolset(
452 maps: &CapabilityMaps,
453 reference: &ToolsetRef,
454) -> Result<ToolRegistry, DistributedCapabilityError> {
455 let registry = maps
456 .toolsets
457 .get(&key(&reference.capability_ref()))
458 .cloned()
459 .ok_or_else(|| unknown("toolset", &reference.capability_ref()))?;
460 let actual = toolset_schema_digest(®istry)?;
461 if actual != reference.schema_digest {
462 return Err(DistributedCapabilityError::new(format!(
463 "toolset {}@{} schema digest mismatch: expected {}, got {actual}",
464 reference.id, reference.version, reference.schema_digest
465 )));
466 }
467 Ok(registry)
468}
469
470fn resolve_tool_policy(
471 maps: &CapabilityMaps,
472 policy: &DistributedToolPolicy,
473) -> Result<ToolPolicy, DistributedCapabilityError> {
474 let approval = match policy.approval.as_str() {
475 "default" => ApprovalPolicy::Default,
476 "never" => ApprovalPolicy::Never,
477 "always" => ApprovalPolicy::Always,
478 "on_request" => ApprovalPolicy::OnRequest,
479 _ => {
480 return Err(DistributedCapabilityError::new(
481 "tool_policy.approval is unsupported",
482 ))
483 }
484 };
485 let can_use_tool = optional(
486 &maps.tool_predicates,
487 "tool_predicate",
488 &policy.predicate_ref,
489 )?;
490 ToolPolicy {
491 allowed_tools: policy.allowed_tools.clone(),
492 disallowed_tools: policy.disallowed_tools.clone(),
493 approval,
494 can_use_tool,
495 denied_side_effects: policy.denied_side_effects.clone(),
496 denied_capability_tags: policy.denied_capability_tags.clone(),
497 deny_terminal_tools: policy.deny_terminal_tools,
498 denied_cost_dimensions: policy.denied_cost_dimensions.clone(),
499 }
500 .normalized()
501 .map_err(|error| DistributedCapabilityError::new(error.to_string()))
502}
503
504fn resolve_checkpoint_extension(
505 maps: &CapabilityMaps,
506 descriptor: &DistributedCheckpointExtensionRef,
507) -> Result<ResolvedDistributedCheckpointExtension, DistributedCapabilityError> {
508 let extension = maps
509 .checkpoint_extensions
510 .get(&key(&descriptor.reference))
511 .cloned()
512 .ok_or_else(|| unknown("checkpoint_extension", &descriptor.reference))?;
513 if extension.namespace() != descriptor.namespace {
514 return Err(DistributedCapabilityError::new(format!(
515 "checkpoint extension {}@{} namespace mismatch: expected {}, got {}",
516 descriptor.reference.id,
517 descriptor.reference.version,
518 descriptor.namespace,
519 extension.namespace()
520 )));
521 }
522 if descriptor.required && !extension.required() {
523 return Err(DistributedCapabilityError::new(format!(
524 "required checkpoint extension {} is registered as optional",
525 descriptor.namespace
526 )));
527 }
528 Ok(ResolvedDistributedCheckpointExtension {
529 descriptor: descriptor.clone(),
530 extension,
531 })
532}
533
534fn optional<T: Clone>(
535 values: &BTreeMap<CapabilityKey, T>,
536 kind: &str,
537 reference: &Option<CapabilityRef>,
538) -> Result<Option<T>, DistributedCapabilityError> {
539 reference
540 .as_ref()
541 .map(|reference| {
542 values
543 .get(&key(reference))
544 .cloned()
545 .ok_or_else(|| unknown(kind, reference))
546 })
547 .transpose()
548}
549
550fn required_many<T: Clone>(
551 values: &BTreeMap<CapabilityKey, T>,
552 kind: &str,
553 references: &[CapabilityRef],
554) -> Result<Vec<T>, DistributedCapabilityError> {
555 references
556 .iter()
557 .map(|reference| {
558 values
559 .get(&key(reference))
560 .cloned()
561 .ok_or_else(|| unknown(kind, reference))
562 })
563 .collect()
564}
565
566fn key(reference: &CapabilityRef) -> CapabilityKey {
567 (reference.id.clone(), reference.version.clone())
568}
569
570fn unknown(kind: &str, reference: &CapabilityRef) -> DistributedCapabilityError {
571 DistributedCapabilityError::new(format!(
572 "unknown distributed capability {kind} {}@{}",
573 reference.id, reference.version
574 ))
575}