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
11pub 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 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 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 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 .is_some_and(|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_or_else(|e| panic!("并发注册任务失败: {e}"));
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_or_else(|e| panic!("并发注册任务失败: {e}"));
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_or_else(|e| panic!("并发注册任务失败: {e}"));
529 }
530 assert_eq!(registry.len(), 1);
531 }
532}