Skip to main content

dtmrs_server/
registry.rs

1//! 进程内分支注册表 —— 嵌入式模式的核心。
2//!
3//! # 为什么分支要用「名字」而不是闭包
4//!
5//! 事务必须能跨进程重启恢复,而**闭包没法持久化**。所以数据库里存的是名字
6//! (`local://deduct`),重启后靠注册表把名字重新解析成函数。
7//!
8//! 这是持久化执行引擎的通用做法,也是唯一正确的做法:
9//!
10//! ```text
11//! 提交时   steps = ["local://deduct", "local://deduct_undo"]  → 落库
12//! 崩溃重启 从库里读出 "local://deduct" → registry 查表 → 拿到函数 → 继续推
13//! ```
14//!
15//! **代价**:注册表在重启后必须注册同样的名字,否则事务推不动。这不是缺陷,
16//! 是把"代码版本"这个隐式依赖显式化了 —— 漏注册会明确报错,而不是静默跑错。
17
18use dtmrs_core::{BranchOp, BranchResult};
19use std::collections::HashMap;
20use std::future::Future;
21use std::pin::Pin;
22use std::sync::Arc;
23
24/// 分支被调用时拿到的上下文。业务侧用它做幂等(配合 dtmrs-barrier)。
25#[derive(Debug, Clone)]
26pub struct BranchCtx {
27    pub gid: String,
28    pub branch_id: String,
29    pub op: BranchOp,
30    pub trans_type: String,
31}
32
33type BoxFut = Pin<Box<dyn Future<Output = BranchResult> + Send>>;
34type Handler = Arc<dyn Fn(BranchCtx) -> BoxFut + Send + Sync>;
35
36/// 分支目标:进程内函数、远端 HTTP,还是远端 gRPC。
37///
38/// 用 URI 前缀区分而不是加新字段 —— 这样落库格式不变,也跟 DTM 的 http URL
39/// 完全兼容,同一个事务里可以三种混用。
40#[derive(Debug, Clone, PartialEq, Eq)]
41pub enum Target {
42    /// `local://名字`
43    Local(String),
44    /// `http://...` / `https://...`
45    Http(String),
46    /// `grpc://host:port/包.服务/方法`(明文)
47    /// 或 `grpcs://host:port/包.服务/方法`(TLS)
48    Grpc(GrpcTarget),
49}
50
51/// 拆好的 gRPC 分支地址。
52///
53/// gRPC 的调用地址天然是两段:连哪个 server(endpoint)+ 调哪个方法(path),
54/// 而 HTTP 是一整个 URL。所以这里必须拆开存,不能像 http 那样原样透传。
55#[derive(Debug, Clone, PartialEq, Eq)]
56pub struct GrpcTarget {
57    /// tonic 连接用。`grpc://` → `http://host:port`,`grpcs://` → `https://host:port`
58    pub endpoint: String,
59    /// gRPC 方法路径,形如 `/包.服务/方法`
60    pub path: String,
61    /// 是否走 TLS。**不能只看 endpoint 的前缀来推**——判定散在两处早晚会漂移,
62    /// 而漂移的后果是「以为加密了其实是明文」,这种错不会有任何报错提示
63    pub tls: bool,
64}
65
66pub fn parse_target(s: &str) -> Target {
67    if let Some(name) = s.strip_prefix("local://") {
68        return Target::Local(name.to_string());
69    }
70    // ⚠ grpcs 必须排在 grpc 前面判断。反过来的话 "grpcs://..." 会先被
71    //   strip_prefix("grpc://") 试探 —— 那个不匹配(因为第 5 个字符是 s),
72    //   所以现在顺序其实无所谓,但写死这个顺序是防止以后有人改成
73    //   starts_with("grpc") 那种前缀判断,那时静默降级成明文就发生了
74    for (prefix, scheme, tls) in [
75        ("grpcs://", "https", true),
76        ("grpc://", "http", false),
77    ] {
78        if let Some(rest) = s.strip_prefix(prefix) {
79            // 认不出来就落到 Http 分支去**明确失败**,不猜。
80            // 静默用错协议比报错难查得多
81            if let Some(t) = parse_grpc(rest, scheme, tls) {
82                return Target::Grpc(t);
83            }
84        }
85    }
86    Target::Http(s.to_string())
87}
88
89/// `host:port/包.服务/方法` → (endpoint, path)
90///
91/// 必须正好有两段路径(服务名 + 方法名)。少一段或多一段都说明地址写错了,
92/// 这时候**不能猜** —— 返回 None 让它落到 Http 分支去明确失败。
93fn parse_grpc(rest: &str, scheme: &str, tls: bool) -> Option<GrpcTarget> {
94    let (authority, path) = rest.split_once('/')?;
95    if authority.is_empty() {
96        return None;
97    }
98    let segs: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect();
99    if segs.len() != 2 {
100        return None;
101    }
102    Some(GrpcTarget {
103        endpoint: format!("{scheme}://{authority}"),
104        path: format!("/{}/{}", segs[0], segs[1]),
105        tls,
106    })
107}
108
109#[derive(Default)]
110pub struct Registry {
111    handlers: HashMap<String, Handler>,
112}
113
114impl Registry {
115    pub fn new() -> Self {
116        Self::default()
117    }
118
119    /// 注册一个进程内分支。名字要跟 `local://名字` 对应。
120    pub fn register<F, Fut>(&mut self, name: &str, f: F) -> &mut Self
121    where
122        F: Fn(BranchCtx) -> Fut + Send + Sync + 'static,
123        Fut: Future<Output = BranchResult> + Send + 'static,
124    {
125        let h: Handler = Arc::new(move |ctx| Box::pin(f(ctx)));
126        self.handlers.insert(name.to_string(), h);
127        self
128    }
129
130    pub fn get(&self, name: &str) -> Option<Handler> {
131        self.handlers.get(name).cloned()
132    }
133
134    pub fn names(&self) -> Vec<&str> {
135        self.handlers.keys().map(String::as_str).collect()
136    }
137
138    /// 提交前自查:所有 `local://` 分支都注册了吗?
139    ///
140    /// 宁可在提交时就报错,也不要等事务推到一半才发现分支不存在 ——
141    /// 那时候已经有副作用落地了,只能靠补偿收拾。
142    pub fn check_all(&self, targets: &[String]) -> Result<(), Vec<String>> {
143        let missing: Vec<String> = targets
144            .iter()
145            .filter_map(|t| match parse_target(t) {
146                Target::Local(n) if !self.handlers.contains_key(&n) => Some(n),
147                _ => None,
148            })
149            .collect();
150        if missing.is_empty() {
151            Ok(())
152        } else {
153            Err(missing)
154        }
155    }
156}
157
158impl std::fmt::Debug for Registry {
159    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
160        f.debug_struct("Registry")
161            .field("handlers", &self.names())
162            .finish()
163    }
164}
165
166#[cfg(test)]
167mod tests {
168    use super::*;
169
170    #[test]
171    fn 前缀区分本地与远端() {
172        assert_eq!(
173            parse_target("local://deduct"),
174            Target::Local("deduct".into())
175        );
176        assert_eq!(
177            parse_target("http://busi/deduct"),
178            Target::Http("http://busi/deduct".into())
179        );
180        // 没前缀就当 http,保持跟 DTM 的兼容
181        assert_eq!(
182            parse_target("https://a/b"),
183            Target::Http("https://a/b".into())
184        );
185    }
186
187    #[test]
188    fn grpc地址拆成端点与方法路径() {
189        let Target::Grpc(t) = parse_target("grpc://127.0.0.1:9000/busi.Busi/Deduct") else {
190            panic!("应该认成 grpc");
191        };
192        // endpoint 必须带 http:// —— tonic 的 Endpoint 要求是个完整 URI
193        assert_eq!(t.endpoint, "http://127.0.0.1:9000");
194        assert_eq!(t.path, "/busi.Busi/Deduct");
195        assert!(!t.tls, "grpc:// 是明文");
196    }
197
198    #[test]
199    fn grpcs走tls且端点是https() {
200        let Target::Grpc(t) = parse_target("grpcs://busi.internal:9000/busi.Busi/Deduct") else {
201            panic!("应该认成 grpc");
202        };
203        assert_eq!(t.endpoint, "https://busi.internal:9000");
204        assert_eq!(t.path, "/busi.Busi/Deduct");
205        assert!(t.tls, "grpcs:// 必须走 TLS");
206    }
207
208    /// ⚠ 这条钉的是**静默降级**:如果 grpcs 因为某种原因没被认出来,
209    /// 它会落到 Http 分支去明确失败 —— 而绝不能变成一个 tls=false 的 Grpc。
210    /// 后者的后果是「以为加密了其实是明文」,没有任何报错提示。
211    #[test]
212    fn grpcs绝不能静默降级成明文() {
213        for s in [
214            "grpcs://a:1/p.S/M",
215            "grpcs://a:1/bad",       // 畸形,会落回 Http
216            "grpcs://",              // 畸形
217        ] {
218            if let Target::Grpc(t) = parse_target(s) {
219                assert!(t.tls, "{s} 认成了 grpc 却没开 TLS —— 这是静默降级成明文");
220                assert!(
221                    t.endpoint.starts_with("https://"),
222                    "{s} 的端点不是 https:{}",
223                    t.endpoint
224                );
225            }
226        }
227    }
228
229    #[test]
230    fn 畸形grpc地址不猜而是落回http() {
231        // 少了方法名、少了服务名、路径多一段、没有 authority ——
232        // 全都不能猜。落到 Http 分支会明确失败,比连错服务安全。
233        for bad in [
234            "grpc://127.0.0.1:9000/onlyservice",
235            "grpc://127.0.0.1:9000/",
236            "grpc://127.0.0.1:9000/a/b/c",
237            "grpc:///a/b",
238            "grpc://noslash",
239        ] {
240            assert!(
241                matches!(parse_target(bad), Target::Http(_)),
242                "{bad} 不该被当成合法 grpc 地址"
243            );
244        }
245    }
246
247    #[tokio::test]
248    async fn 注册与调用() {
249        let mut r = Registry::new();
250        r.register("ok", |_ctx| async { BranchResult::Success });
251        let h = r.get("ok").expect("应该能查到");
252        let ctx = BranchCtx {
253            gid: "g".into(),
254            branch_id: "01".into(),
255            op: BranchOp::Action,
256            trans_type: "saga".into(),
257        };
258        assert_eq!(h(ctx).await, BranchResult::Success);
259        assert!(r.get("nope").is_none());
260    }
261
262    #[test]
263    fn 提交前能查出漏注册的分支() {
264        let mut r = Registry::new();
265        r.register("a", |_| async { BranchResult::Success });
266        let targets = vec![
267            "local://a".to_string(),
268            "local://missing".to_string(),
269            "http://x/y".to_string(),
270        ];
271        let err = r.check_all(&targets).unwrap_err();
272        assert_eq!(err, vec!["missing"], "只报本地漏的,http 不管");
273        assert!(r.check_all(&["local://a".to_string()]).is_ok());
274    }
275}