Skip to main content

sz_rust_capability/
registry.rs

1use std::collections::HashMap;
2use std::sync::atomic::{AtomicU64, Ordering};
3use std::sync::Arc;
4
5use crate::capability::{Capability, CapabilityInfo};
6use crate::error::{CapError, CapResult};
7use crate::metrics::CapMetrics;
8use crate::permission::PermissionChecker;
9use crate::source::CapabilitySource;
10
11/// 中心能力注册表,提供注册/发现/调用的统一入口。
12///
13/// 内部使用 `parking_lot::RwLock<HashMap<String, Arc<dyn Capability>>>` 保证并发安全。
14/// 所有读操作(get/find/search/list)在读锁内完成 Arc 克隆后释放锁,不跨 await 点。
15///
16/// # 调用链路
17///
18/// `call` / `call_with_tenant` 执行三步链路:参数校验 → 权限检查 → 能力调用。
19/// 未设置 `PermissionChecker` 时默认放行。
20///
21/// # 性能指标
22///
23/// | 操作 | 延迟 |
24/// |------|------|
25/// | register | ~187 ns |
26/// | get | ~38 ns |
27/// | find_by_tags (1000 能力) | ~20 μs |
28pub struct CapabilityRegistry {
29    capabilities: parking_lot::RwLock<HashMap<String, Arc<dyn Capability>>>,
30    permission_checker: parking_lot::RwLock<Option<Arc<dyn PermissionChecker>>>,
31    call_total: AtomicU64,
32}
33
34impl CapabilityRegistry {
35    pub fn new() -> Self {
36        Self {
37            capabilities: parking_lot::RwLock::new(HashMap::new()),
38            permission_checker: parking_lot::RwLock::new(None),
39            call_total: AtomicU64::new(0),
40        }
41    }
42
43    /// 设置权限检查器,设置后所有 `call` / `call_with_tenant` 将在能力调用前执行权限检查。
44    pub fn set_permission_checker(&self, checker: Arc<dyn PermissionChecker>) {
45        let mut guard = self.permission_checker.write();
46        *guard = Some(checker);
47    }
48
49    pub fn register(&self, cap: Arc<dyn Capability>) -> Option<Arc<dyn Capability>> {
50        let name = cap.name().to_string();
51        let mut caps = self.capabilities.write();
52        caps.insert(name, cap)
53    }
54
55    pub fn unregister(&self, name: &str) -> Option<Arc<dyn Capability>> {
56        let mut caps = self.capabilities.write();
57        caps.remove(name)
58    }
59
60    pub fn get(&self, name: &str) -> Option<Arc<dyn Capability>> {
61        let caps = self.capabilities.read();
62        caps.get(name).cloned()
63    }
64
65    pub fn find_by_tags(
66        &self,
67        tags: &[&str],
68        source: Option<CapabilitySource>,
69    ) -> Vec<Arc<dyn Capability>> {
70        let caps = self.capabilities.read();
71        caps.values()
72            .filter(|cap| {
73                let tag_match = tags.iter().all(|t| cap.tags().contains(t));
74                let source_match = source.map_or(true, |s| cap.source() == s);
75                tag_match && source_match
76            })
77            .cloned()
78            .collect()
79    }
80
81    pub fn search(&self, query: &str) -> Vec<Arc<dyn Capability>> {
82        let caps = self.capabilities.read();
83        let query_lower = query.to_lowercase();
84        caps.values()
85            .filter(|cap| {
86                cap.name().to_lowercase().contains(&query_lower)
87                    || cap.description().to_lowercase().contains(&query_lower)
88            })
89            .cloned()
90            .collect()
91    }
92
93    pub fn list_all(&self) -> Vec<Arc<dyn Capability>> {
94        let caps = self.capabilities.read();
95        caps.values().cloned().collect()
96    }
97
98    pub fn list_by_source(&self, source: CapabilitySource) -> Vec<Arc<dyn Capability>> {
99        let caps = self.capabilities.read();
100        caps.values()
101            .filter(|cap| cap.source() == source)
102            .cloned()
103            .collect()
104    }
105
106    /// 调用能力,使用默认 `tenant_id = 0`(无租户上下文)。
107    ///
108    /// 执行链路:参数校验 → 权限检查 → 能力调用。
109    pub async fn call(&self, name: &str, args: serde_json::Value) -> CapResult<serde_json::Value> {
110        self.call_with_tenant(name, args, 0).await
111    }
112
113    /// 调用能力,携带租户上下文用于权限检查。
114    ///
115    /// 执行链路:参数校验 → 权限检查 → 能力调用。
116    /// 未设置 `PermissionChecker` 时默认放行。
117    pub async fn call_with_tenant(
118        &self,
119        name: &str,
120        args: serde_json::Value,
121        tenant_id: i64,
122    ) -> CapResult<serde_json::Value> {
123        let cap = self
124            .get(name)
125            .ok_or_else(|| CapError::NotFound(name.to_string()))?;
126
127        if cap.requires_confirmation() {
128            return Err(CapError::ConfirmationRequired);
129        }
130
131        cap.validate_args(&args).await?;
132
133        let checker_opt = {
134            let guard = self.permission_checker.read();
135            guard.as_ref().cloned()
136        };
137        if let Some(checker) = checker_opt {
138            checker.check(name, &args, tenant_id).await?;
139        }
140
141        self.call_total.fetch_add(1, Ordering::Relaxed);
142        cap.call(args).await
143    }
144
145    pub fn len(&self) -> usize {
146        let caps = self.capabilities.read();
147        caps.len()
148    }
149
150    pub fn is_empty(&self) -> bool {
151        self.len() == 0
152    }
153
154    pub fn metrics(&self) -> CapMetrics {
155        let caps = self.capabilities.read();
156        let mut by_source = HashMap::new();
157        for cap in caps.values() {
158            *by_source.entry(cap.source()).or_insert(0) += 1;
159        }
160        CapMetrics {
161            total: caps.len(),
162            by_source,
163            call_total: self.call_total.load(Ordering::Relaxed),
164        }
165    }
166
167    pub fn list_info(&self) -> Vec<CapabilityInfo> {
168        let caps = self.capabilities.read();
169        caps.values()
170            .map(|cap| CapabilityInfo::from_trait(cap.as_ref()))
171            .collect()
172    }
173}
174
175impl Default for CapabilityRegistry {
176    fn default() -> Self {
177        Self::new()
178    }
179}
180
181pub fn validate_json_schema(schema: &serde_json::Value, args: &serde_json::Value) -> CapResult<()> {
182    if !schema.is_object() {
183        return Ok(());
184    }
185
186    if let Some(required) = schema.get("required").and_then(|v| v.as_array()) {
187        for field in required {
188            if let Some(field_name) = field.as_str() {
189                if !args
190                    .as_object()
191                    .map_or(false, |obj| obj.contains_key(field_name))
192                {
193                    return Err(CapError::ValidationError(format!(
194                        "缺少必填字段: {field_name}"
195                    )));
196                }
197            }
198        }
199    }
200
201    if let (Some(properties), Some(args_obj)) = (schema.get("properties"), args.as_object()) {
202        if let Some(props) = properties.as_object() {
203            for (field_name, field_schema) in props {
204                if let Some(field_value) = args_obj.get(field_name) {
205                    if let Some(expected_type) = field_schema.get("type").and_then(|v| v.as_str()) {
206                        let actual_type = json_type_of(field_value);
207                        let type_match = expected_type == actual_type
208                            || (expected_type == "number" && actual_type == "integer");
209                        if !type_match {
210                            return Err(CapError::ValidationError(format!(
211                                "字段 {field_name} 类型不匹配: 期望 {expected_type}, 实际 {actual_type}"
212                            )));
213                        }
214                    }
215                }
216            }
217        }
218    }
219
220    Ok(())
221}
222
223fn json_type_of(value: &serde_json::Value) -> &'static str {
224    match value {
225        serde_json::Value::Null => "null",
226        serde_json::Value::Bool(_) => "boolean",
227        serde_json::Value::Number(n) => {
228            if n.is_i64() || n.is_u64() {
229                "integer"
230            } else {
231                "number"
232            }
233        }
234        serde_json::Value::String(_) => "string",
235        serde_json::Value::Array(_) => "array",
236        serde_json::Value::Object(_) => "object",
237    }
238}
239
240#[cfg(test)]
241mod tests {
242    use super::*;
243    use async_trait::async_trait;
244    use serde_json::json;
245
246    struct EchoCapability;
247
248    #[async_trait]
249    impl Capability for EchoCapability {
250        fn name(&self) -> &'static str {
251            "echo"
252        }
253        fn description(&self) -> &'static str {
254            "回显输入参数"
255        }
256        fn schema(&self) -> serde_json::Value {
257            json!({
258                "type": "object",
259                "properties": {
260                    "message": { "type": "string" }
261                },
262                "required": ["message"]
263            })
264        }
265        fn tags(&self) -> &[&'static str] {
266            &["test", "echo"]
267        }
268        fn source(&self) -> CapabilitySource {
269            CapabilitySource::Skill
270        }
271        async fn call(&self, args: serde_json::Value) -> CapResult<serde_json::Value> {
272            Ok(args)
273        }
274    }
275
276    #[tokio::test]
277    async fn test_register_and_get() {
278        let registry = CapabilityRegistry::new();
279        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
280        let old = registry.register(cap);
281        assert!(old.is_none());
282        assert!(registry.get("echo").is_some());
283        assert_eq!(registry.len(), 1);
284    }
285
286    #[tokio::test]
287    async fn test_register_overwrite() {
288        let registry = CapabilityRegistry::new();
289        let cap1 = Arc::new(EchoCapability) as Arc<dyn Capability>;
290        let cap2 = Arc::new(EchoCapability) as Arc<dyn Capability>;
291        registry.register(cap1);
292        let old = registry.register(cap2);
293        assert!(old.is_some());
294        assert_eq!(registry.len(), 1);
295    }
296
297    #[tokio::test]
298    async fn test_unregister() {
299        let registry = CapabilityRegistry::new();
300        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
301        registry.register(cap);
302        let removed = registry.unregister("echo");
303        assert!(removed.is_some());
304        assert!(registry.get("echo").is_none());
305        assert!(registry.is_empty());
306    }
307
308    #[tokio::test]
309    async fn test_find_by_tags() {
310        let registry = CapabilityRegistry::new();
311        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
312        registry.register(cap);
313        let caps = registry.find_by_tags(&["test"], None);
314        assert_eq!(caps.len(), 1);
315        let caps = registry.find_by_tags(&["test"], Some(CapabilitySource::Skill));
316        assert_eq!(caps.len(), 1);
317        let caps = registry.find_by_tags(&["test"], Some(CapabilitySource::Plugin));
318        assert_eq!(caps.len(), 0);
319        let caps = registry.find_by_tags(&["nonexistent"], None);
320        assert_eq!(caps.len(), 0);
321    }
322
323    #[tokio::test]
324    async fn test_search() {
325        let registry = CapabilityRegistry::new();
326        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
327        registry.register(cap);
328        let caps = registry.search("echo");
329        assert_eq!(caps.len(), 1);
330        let caps = registry.search("回显");
331        assert_eq!(caps.len(), 1);
332        let caps = registry.search("nonexistent");
333        assert_eq!(caps.len(), 0);
334    }
335
336    #[tokio::test]
337    async fn test_call_success() {
338        let registry = CapabilityRegistry::new();
339        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
340        registry.register(cap);
341        let result = registry.call("echo", json!({"message": "hello"})).await;
342        assert!(result.is_ok());
343        assert_eq!(result.unwrap(), json!({"message": "hello"}));
344    }
345
346    #[tokio::test]
347    async fn test_call_not_found() {
348        let registry = CapabilityRegistry::new();
349        let result = registry.call("nonexistent", json!({})).await;
350        assert!(matches!(result, Err(CapError::NotFound(_))));
351    }
352
353    #[tokio::test]
354    async fn test_call_validation_error() {
355        let registry = CapabilityRegistry::new();
356        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
357        registry.register(cap);
358        let result = registry.call("echo", json!({})).await;
359        assert!(matches!(result, Err(CapError::ValidationError(_))));
360    }
361
362    #[tokio::test]
363    async fn test_call_type_mismatch() {
364        let registry = CapabilityRegistry::new();
365        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
366        registry.register(cap);
367        let result = registry.call("echo", json!({"message": 123})).await;
368        assert!(matches!(result, Err(CapError::ValidationError(_))));
369    }
370
371    #[tokio::test]
372    async fn test_metrics() {
373        let registry = CapabilityRegistry::new();
374        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
375        registry.register(cap);
376        let _ = registry.call("echo", json!({"message": "hi"})).await;
377        let metrics = registry.metrics();
378        assert_eq!(metrics.total, 1);
379        assert_eq!(metrics.call_total, 1);
380        assert_eq!(metrics.by_source.get(&CapabilitySource::Skill), Some(&1));
381    }
382
383    #[test]
384    fn test_validate_json_schema_valid() {
385        let schema = json!({
386            "type": "object",
387            "properties": { "name": { "type": "string" } },
388            "required": ["name"]
389        });
390        let args = json!({ "name": "test" });
391        assert!(validate_json_schema(&schema, &args).is_ok());
392    }
393
394    #[test]
395    fn test_validate_json_schema_missing_required() {
396        let schema = json!({
397            "required": ["name"]
398        });
399        let args = json!({});
400        assert!(validate_json_schema(&schema, &args).is_err());
401    }
402
403    #[test]
404    fn test_validate_json_schema_type_mismatch() {
405        let schema = json!({
406            "properties": { "age": { "type": "number" } }
407        });
408        let args = json!({ "age": "twenty" });
409        assert!(validate_json_schema(&schema, &args).is_err());
410    }
411
412    #[test]
413    fn test_validate_json_schema_no_schema() {
414        let args = json!({ "any": "thing" });
415        assert!(validate_json_schema(&json!(null), &args).is_ok());
416    }
417
418    #[tokio::test]
419    async fn test_concurrent_register_and_get() {
420        use std::sync::Arc;
421        let registry = Arc::new(CapabilityRegistry::new());
422
423        struct NamedCap {
424            cap_name: &'static str,
425        }
426        #[async_trait]
427        impl Capability for NamedCap {
428            fn name(&self) -> &'static str {
429                self.cap_name
430            }
431            fn description(&self) -> &'static str {
432                "并发测试能力"
433            }
434            fn schema(&self) -> serde_json::Value {
435                json!({})
436            }
437            fn tags(&self) -> &[&'static str] {
438                &["concurrent"]
439            }
440            fn source(&self) -> CapabilitySource {
441                CapabilitySource::Skill
442            }
443            async fn call(&self, args: serde_json::Value) -> CapResult<serde_json::Value> {
444                Ok(args)
445            }
446        }
447
448        let mut handles = vec![];
449        for i in 0..50u32 {
450            let reg = registry.clone();
451            handles.push(tokio::spawn(async move {
452                let name: &'static str = Box::leak(format!("cap_{i}").into_boxed_str());
453                let cap = Arc::new(NamedCap { cap_name: name }) as Arc<dyn Capability>;
454                reg.register(cap);
455            }));
456        }
457        for h in handles {
458            h.await.unwrap();
459        }
460        assert_eq!(registry.len(), 50);
461
462        let caps = registry.find_by_tags(&["concurrent"], None);
463        assert_eq!(caps.len(), 50);
464    }
465
466    #[tokio::test]
467    async fn test_concurrent_call_no_deadlock() {
468        use std::sync::Arc;
469        let registry = Arc::new(CapabilityRegistry::new());
470        let cap = Arc::new(EchoCapability) as Arc<dyn Capability>;
471        registry.register(cap);
472
473        let mut handles = vec![];
474        for _ in 0..100u32 {
475            let reg = registry.clone();
476            handles.push(tokio::spawn(async move {
477                let _ = reg.call("echo", json!({"message": "concurrent"})).await;
478            }));
479        }
480        for h in handles {
481            h.await.unwrap();
482        }
483        assert_eq!(registry.metrics().call_total, 100);
484    }
485
486    #[tokio::test]
487    async fn test_concurrent_mixed_operations() {
488        use std::sync::Arc;
489        let registry = Arc::new(CapabilityRegistry::new());
490
491        struct MixedCap;
492        #[async_trait]
493        impl Capability for MixedCap {
494            fn name(&self) -> &'static str {
495                "mixed_cap"
496            }
497            fn description(&self) -> &'static str {
498                "混合操作测试"
499            }
500            fn schema(&self) -> serde_json::Value {
501                json!({})
502            }
503            fn tags(&self) -> &[&'static str] {
504                &["mixed"]
505            }
506            fn source(&self) -> CapabilitySource {
507                CapabilitySource::Plugin
508            }
509            async fn call(&self, args: serde_json::Value) -> CapResult<serde_json::Value> {
510                Ok(args)
511            }
512        }
513
514        let cap = Arc::new(MixedCap) as Arc<dyn Capability>;
515        registry.register(cap);
516
517        let mut handles = vec![];
518        for _ in 0..50 {
519            let reg = registry.clone();
520            handles.push(tokio::spawn(async move {
521                let _ = reg.get("mixed_cap");
522                let _ = reg.find_by_tags(&["mixed"], None);
523                let _ = reg.list_all();
524                let _ = reg.metrics();
525            }));
526        }
527        for h in handles {
528            h.await.unwrap();
529        }
530        assert_eq!(registry.len(), 1);
531    }
532}