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