use dtmrs_core::{BranchOp, BranchResult};
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct BranchCtx {
pub gid: String,
pub branch_id: String,
pub op: BranchOp,
pub trans_type: String,
}
type BoxFut = Pin<Box<dyn Future<Output = BranchResult> + Send>>;
type Handler = Arc<dyn Fn(BranchCtx) -> BoxFut + Send + Sync>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Target {
Local(String),
Http(String),
Grpc(GrpcTarget),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GrpcTarget {
pub endpoint: String,
pub path: String,
}
pub fn parse_target(s: &str) -> Target {
if let Some(name) = s.strip_prefix("local://") {
return Target::Local(name.to_string());
}
if let Some(rest) = s.strip_prefix("grpc://") {
if let Some(t) = parse_grpc(rest) {
return Target::Grpc(t);
}
}
Target::Http(s.to_string())
}
fn parse_grpc(rest: &str) -> Option<GrpcTarget> {
let (authority, path) = rest.split_once('/')?;
if authority.is_empty() {
return None;
}
let segs: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect();
if segs.len() != 2 {
return None;
}
Some(GrpcTarget {
endpoint: format!("http://{authority}"),
path: format!("/{}/{}", segs[0], segs[1]),
})
}
#[derive(Default)]
pub struct Registry {
handlers: HashMap<String, Handler>,
}
impl Registry {
pub fn new() -> Self {
Self::default()
}
pub fn register<F, Fut>(&mut self, name: &str, f: F) -> &mut Self
where
F: Fn(BranchCtx) -> Fut + Send + Sync + 'static,
Fut: Future<Output = BranchResult> + Send + 'static,
{
let h: Handler = Arc::new(move |ctx| Box::pin(f(ctx)));
self.handlers.insert(name.to_string(), h);
self
}
pub fn get(&self, name: &str) -> Option<Handler> {
self.handlers.get(name).cloned()
}
pub fn names(&self) -> Vec<&str> {
self.handlers.keys().map(String::as_str).collect()
}
pub fn check_all(&self, targets: &[String]) -> Result<(), Vec<String>> {
let missing: Vec<String> = targets
.iter()
.filter_map(|t| match parse_target(t) {
Target::Local(n) if !self.handlers.contains_key(&n) => Some(n),
_ => None,
})
.collect();
if missing.is_empty() {
Ok(())
} else {
Err(missing)
}
}
}
impl std::fmt::Debug for Registry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Registry")
.field("handlers", &self.names())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn 前缀区分本地与远端() {
assert_eq!(
parse_target("local://deduct"),
Target::Local("deduct".into())
);
assert_eq!(
parse_target("http://busi/deduct"),
Target::Http("http://busi/deduct".into())
);
assert_eq!(
parse_target("https://a/b"),
Target::Http("https://a/b".into())
);
}
#[test]
fn grpc地址拆成端点与方法路径() {
let Target::Grpc(t) = parse_target("grpc://127.0.0.1:9000/busi.Busi/Deduct") else {
panic!("应该认成 grpc");
};
assert_eq!(t.endpoint, "http://127.0.0.1:9000");
assert_eq!(t.path, "/busi.Busi/Deduct");
}
#[test]
fn 畸形grpc地址不猜而是落回http() {
for bad in [
"grpc://127.0.0.1:9000/onlyservice",
"grpc://127.0.0.1:9000/",
"grpc://127.0.0.1:9000/a/b/c",
"grpc:///a/b",
"grpc://noslash",
] {
assert!(
matches!(parse_target(bad), Target::Http(_)),
"{bad} 不该被当成合法 grpc 地址"
);
}
}
#[tokio::test]
async fn 注册与调用() {
let mut r = Registry::new();
r.register("ok", |_ctx| async { BranchResult::Success });
let h = r.get("ok").expect("应该能查到");
let ctx = BranchCtx {
gid: "g".into(),
branch_id: "01".into(),
op: BranchOp::Action,
trans_type: "saga".into(),
};
assert_eq!(h(ctx).await, BranchResult::Success);
assert!(r.get("nope").is_none());
}
#[test]
fn 提交前能查出漏注册的分支() {
let mut r = Registry::new();
r.register("a", |_| async { BranchResult::Success });
let targets = vec![
"local://a".to_string(),
"local://missing".to_string(),
"http://x/y".to_string(),
];
let err = r.check_all(&targets).unwrap_err();
assert_eq!(err, vec!["missing"], "只报本地漏的,http 不管");
assert!(r.check_all(&["local://a".to_string()]).is_ok());
}
}