Skip to main content

dtmrs_server/
http.rs

1//! HTTP 协议层:把 HTTP 请求翻译成 `Api` 调用,再把结果翻译成 DTM 的应答格式。
2//!
3//! ⚠ **这一层只做协议转换,不做任何业务判断** —— 判断全在 `api.rs`。
4//! gRPC 层(`grpc/server.rs`)是它的对偶,两边必须对同一个请求给出同样的受理/
5//! 拒绝结论。否则会出现「同一个请求走 HTTP 被拒、走 gRPC 却受理了」。
6//!
7//! 这个模块**刻意放在 dtmrs-server 而不是二进制 crate 里**:早先它写在
8//! `main.rs`,测试够不着,覆盖率是 0%,而 gRPC 层有 86% —— 防漂移的约束
9//! 只有一半受测试保护。搬过来之后 `router()` 可导出,两边就能用同一组
10//! 用例做等价性测试(见 tests/http.rs 的「HTTP 与 gRPC 等价」那几个)。
11
12use crate::api::{Api, ApiError, RegisterBranch, TransView};
13use axum::extract::{Query, State};
14use axum::http::StatusCode;
15use axum::routing::{get, post};
16use axum::{Json, Router};
17use dtmrs_core::SagaStep;
18use serde::{Deserialize, Serialize};
19use std::collections::HashMap;
20
21/// axum 的 `State` 要求 Clone;`Api` 内部是 Arc,克隆很便宜。
22#[derive(Clone)]
23pub struct App {
24    api: Api,
25}
26
27impl App {
28    pub fn new(api: Api) -> Self {
29        Self { api }
30    }
31}
32
33#[derive(Deserialize)]
34struct SubmitReq {
35    gid: String,
36    #[serde(default = "default_trans_type")]
37    trans_type: String,
38    /// saga 一次性给全部步骤;tcc/msg 走 prepare + submit,这里可以不带
39    #[serde(default)]
40    steps: Vec<SagaStep>,
41}
42
43/// 二阶段消息 / TCC 的第一阶段
44#[derive(Deserialize)]
45struct PrepareReq {
46    gid: String,
47    trans_type: String,
48    /// msg 用:正向分支列表(没有补偿)
49    #[serde(default)]
50    actions: Vec<String>,
51    /// msg 用:回查地址。进程在 prepare 和 submit 之间崩了,TC 靠它决断
52    #[serde(default)]
53    query_prepared: String,
54    /// msg 用:回查前的宽限秒数,默认 10
55    #[serde(default)]
56    grace_secs: Option<i64>,
57}
58
59/// 分支登记。TCC 用 confirm/cancel,XA 用 commit/rollback。
60#[derive(Deserialize)]
61struct RegisterBranchReq {
62    gid: String,
63    branch_id: String,
64    #[serde(default)]
65    confirm: String,
66    #[serde(default)]
67    cancel: String,
68    /// TCC 的 try,可选,只为可观测性存一份
69    #[serde(default)]
70    r#try: String,
71    #[serde(default)]
72    commit: String,
73    #[serde(default)]
74    rollback: String,
75}
76
77fn default_trans_type() -> String {
78    "saga".into()
79}
80
81#[derive(Serialize)]
82struct Reply {
83    dtm_result: &'static str,
84    #[serde(skip_serializing_if = "Option::is_none")]
85    message: Option<String>,
86}
87
88impl Reply {
89    fn ok() -> Json<Self> {
90        Json(Self {
91            dtm_result: "SUCCESS",
92            message: None,
93        })
94    }
95    fn err(m: impl Into<String>) -> Json<Self> {
96        Json(Self {
97            dtm_result: "FAILURE",
98            message: Some(m.into()),
99        })
100    }
101}
102
103/// [`ApiError`] → HTTP。
104///
105/// `Conflict` 返回 **200 + FAILURE 体**是刻意保留的历史行为(已终结的事务
106/// 再调 abort),换成 4xx 会打破现有客户端。
107fn http_err(e: ApiError) -> (StatusCode, Json<Reply>) {
108    let code = match &e {
109        ApiError::BadRequest(_) => StatusCode::BAD_REQUEST,
110        ApiError::NotFound(_) => StatusCode::NOT_FOUND,
111        ApiError::Conflict(_) => StatusCode::OK,
112        ApiError::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR,
113    };
114    (code, Reply::err(e.message().to_string()))
115}
116
117fn http_result(r: Result<(), ApiError>) -> (StatusCode, Json<Reply>) {
118    match r {
119        Ok(()) => (StatusCode::OK, Reply::ok()),
120        Err(e) => http_err(e),
121    }
122}
123
124
125pub fn router(app: App) -> Router {
126    Router::new()
127        .route("/api/dtmsvr/newGid", get(new_gid))
128        .route("/api/dtmsvr/prepare", post(prepare))
129        .route("/api/dtmsvr/registerBranch", post(register_branch))
130        .route("/api/dtmsvr/submit", post(submit))
131        .route("/api/dtmsvr/abort", post(abort))
132        .route("/api/dtmsvr/retry", post(retry))
133        .route("/api/dtmsvr/query", get(query))
134        .route("/api/dtmsvr/all", get(all))
135        .route("/health", get(|| async { "ok" }))
136        // 管理台。单文件内嵌,没有构建步骤也没有外部依赖 ——
137        // 内网和离线环境都能直接用
138        .route("/", get(console))
139        .route("/console", get(console))
140        .with_state(app)
141}
142
143async fn new_gid(State(app): State<App>) -> Json<HashMap<&'static str, String>> {
144    Json(HashMap::from([("gid", app.api.new_gid())]))
145}
146
147async fn submit(State(app): State<App>, Json(req): Json<SubmitReq>) -> (StatusCode, Json<Reply>) {
148    http_result(app.api.submit(&req.gid, &req.trans_type, &req.steps).await)
149}
150
151async fn prepare(State(app): State<App>, Json(req): Json<PrepareReq>) -> (StatusCode, Json<Reply>) {
152    http_result(
153        app.api
154            .prepare(
155                &req.gid,
156                &req.trans_type,
157                &req.actions,
158                &req.query_prepared,
159                req.grace_secs,
160            )
161            .await,
162    )
163}
164
165async fn register_branch(
166    State(app): State<App>,
167    Json(req): Json<RegisterBranchReq>,
168) -> (StatusCode, Json<Reply>) {
169    http_result(
170        app.api
171            .register_branch(&RegisterBranch {
172                gid: req.gid,
173                branch_id: req.branch_id,
174                confirm: req.confirm,
175                cancel: req.cancel,
176                r#try: req.r#try,
177                commit: req.commit,
178                rollback: req.rollback,
179            })
180            .await,
181    )
182}
183
184#[derive(Deserialize)]
185struct GidQuery {
186    gid: String,
187}
188
189async fn abort(State(app): State<App>, Json(q): Json<GidQuery>) -> (StatusCode, Json<Reply>) {
190    http_result(app.api.abort(&q.gid).await)
191}
192
193/// 立刻重试:把事务排到调度队首。管理台用,也可以直接调
194async fn retry(State(app): State<App>, Json(q): Json<GidQuery>) -> (StatusCode, Json<Reply>) {
195    http_result(app.api.retry(&q.gid).await)
196}
197
198/// 管理台页面。`include_str!` 编进二进制,部署时不用带额外文件
199async fn console() -> axum::response::Html<&'static str> {
200    axum::response::Html(include_str!("console.html"))
201}
202
203async fn query(
204    State(app): State<App>,
205    Query(q): Query<GidQuery>,
206) -> Result<Json<TransView>, (StatusCode, Json<Reply>)> {
207    app.api.query(&q.gid).await.map(Json).map_err(http_err)
208}
209
210async fn all(State(app): State<App>) -> Json<Vec<TransView>> {
211    Json(app.api.list_recent(100).await)
212}