sz_rust_core/plugin/
cross_query.rs1use serde_json::Value;
7
8#[derive(Debug, thiserror::Error)]
10pub enum CrossQueryError {
11 #[error("权限不足:租户 {tenant_id} 无权查询")]
13 PermissionDenied {
14 tenant_id: i64,
16 },
17 #[error("能力未找到: {0}")]
19 NotFound(String),
20 #[error("查询失败: {0}")]
22 QueryFailed(String),
23}
24
25pub struct CrossQuery {
30 tenant_id: i64,
32}
33
34impl CrossQuery {
35 pub fn new(tenant_id: i64) -> Self {
37 Self { tenant_id }
38 }
39
40 pub fn tenant_id(&self) -> i64 {
42 self.tenant_id
43 }
44
45 pub fn inject_tenant_filter(&self, mut args: Value) -> Value {
49 if let Some(obj) = args.as_object_mut() {
50 obj.insert("tenant_id".to_string(), Value::from(self.tenant_id));
51 } else if args.is_null() {
52 args = serde_json::json!({ "tenant_id": self.tenant_id });
53 }
54 args
55 }
56
57 pub fn verify_tenant(&self, target_tenant_id: i64) -> Result<(), CrossQueryError> {
61 if self.tenant_id != target_tenant_id {
62 return Err(CrossQueryError::PermissionDenied {
63 tenant_id: self.tenant_id,
64 });
65 }
66 Ok(())
67 }
68
69 pub fn aggregate(&self, queries: &[(&str, Value)]) -> Value {
74 let queries_json: Vec<Value> = queries
75 .iter()
76 .map(|(name, args)| {
77 let injected = self.inject_tenant_filter(args.clone());
78 serde_json::json!({
79 "capability": name,
80 "args": injected,
81 })
82 })
83 .collect();
84 serde_json::json!({
85 "tenant_id": self.tenant_id,
86 "queries": queries_json,
87 })
88 }
89}
90
91#[cfg(test)]
92mod tests {
93 use super::*;
94
95 #[test]
96 fn test_inject_tenant_filter() {
97 let cq = CrossQuery::new(100);
98 let args = serde_json::json!({"keyword": "test"});
99 let result = cq.inject_tenant_filter(args);
100 assert_eq!(result["tenant_id"], 100);
101 assert_eq!(result["keyword"], "test");
102 }
103
104 #[test]
105 fn test_inject_tenant_filter_null() {
106 let cq = CrossQuery::new(200);
107 let result = cq.inject_tenant_filter(Value::Null);
108 assert_eq!(result["tenant_id"], 200);
109 }
110
111 #[test]
112 fn test_verify_tenant_same() {
113 let cq = CrossQuery::new(100);
114 assert!(cq.verify_tenant(100).is_ok());
115 }
116
117 #[test]
118 fn test_verify_tenant_different() {
119 let cq = CrossQuery::new(100);
120 assert!(cq.verify_tenant(200).is_err());
121 }
122
123 #[test]
124 fn test_aggregate() {
125 let cq = CrossQuery::new(100);
126 let queries = vec![
127 ("plugin_a.search", serde_json::json!({"q": "hello"})),
128 ("plugin_b.list", serde_json::json!({})),
129 ];
130 let result = cq.aggregate(&queries);
131 assert_eq!(result["tenant_id"], 100);
132 assert_eq!(result["queries"].as_array().unwrap().len(), 2);
133 assert_eq!(result["queries"][0]["args"]["tenant_id"], 100);
134 }
135}