Skip to main content

vv_agent/runtime/backends/distributed/
capabilities.rs

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(&registry)?;
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(&registry)?;
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}