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