Files
ystp/src/api/billing.rs

736 lines
22 KiB
Rust

use crate::api::context;
use crate::api::envelope::Envelope;
use crate::error::{AppError, ErrorCode};
use crate::services::billing;
use crate::services::idempotency;
use crate::services::quota;
use crate::services::settings;
use crate::state::AppState;
use axum::extract::{ConnectInfo, State};
use axum::http::HeaderMap;
use axum::routing::{get, post};
use axum::{Json, Router};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use sqlx::FromRow;
use std::net::SocketAddr;
use uuid::Uuid;
pub fn router() -> Router<AppState> {
Router::new()
.route("/billing/plans", get(list_plans))
.route("/billing/subscription", get(get_subscription))
.route("/billing/usage", get(get_usage))
.route("/billing/invoices", get(list_invoices))
.route("/billing/checkout", post(create_checkout))
.route("/billing/portal", post(create_portal))
}
#[derive(Debug, FromRow, Serialize)]
struct PlanView {
id: Uuid,
code: String,
name: String,
currency: String,
amount_cents: i32,
interval: String,
included_units_per_period: i32,
max_file_size_mb: i32,
max_files_per_batch: i32,
retention_days: i32,
features: serde_json::Value,
}
#[derive(Debug, Serialize)]
struct PlansResponse {
plans: Vec<PlanView>,
}
async fn list_plans(
State(state): State<AppState>,
) -> Result<Json<Envelope<PlansResponse>>, AppError> {
let plans = sqlx::query_as::<_, PlanView>(
r#"
SELECT
id, code, name, currency, amount_cents, interval,
included_units_per_period, max_file_size_mb, max_files_per_batch, retention_days,
features
FROM plans
WHERE is_active = true
ORDER BY amount_cents ASC
"#,
)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询套餐失败").with_source(err))?;
Ok(Json(Envelope {
success: true,
data: PlansResponse { plans },
}))
}
#[derive(Debug, Serialize)]
struct SubscriptionPlanView {
id: Uuid,
code: String,
name: String,
currency: String,
amount_cents: i32,
interval: String,
included_units_per_period: i32,
max_file_size_mb: i32,
max_files_per_batch: i32,
retention_days: i32,
features: serde_json::Value,
}
#[derive(Debug, Serialize)]
struct SubscriptionView {
status: String,
current_period_start: DateTime<Utc>,
current_period_end: DateTime<Utc>,
cancel_at_period_end: bool,
plan: SubscriptionPlanView,
}
#[derive(Debug, Serialize)]
struct SubscriptionResponse {
subscription: SubscriptionView,
}
async fn get_subscription(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
) -> Result<Json<Envelope<SubscriptionResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
#[derive(Debug, FromRow)]
struct SubRow {
status: String,
current_period_start: DateTime<Utc>,
current_period_end: DateTime<Utc>,
cancel_at_period_end: bool,
plan_id: Uuid,
plan_code: String,
plan_name: String,
currency: String,
amount_cents: i32,
interval: String,
included_units_per_period: i32,
max_file_size_mb: i32,
max_files_per_batch: i32,
retention_days: i32,
features: serde_json::Value,
}
let sub = sqlx::query_as::<_, SubRow>(
r#"
SELECT
s.status::text AS status,
s.current_period_start,
s.current_period_end,
s.cancel_at_period_end,
p.id AS plan_id,
p.code AS plan_code,
p.name AS plan_name,
p.currency,
p.amount_cents,
p.interval,
p.included_units_per_period,
p.max_file_size_mb,
p.max_files_per_batch,
p.retention_days,
p.features
FROM subscriptions s
JOIN plans p ON p.id = s.plan_id
WHERE s.user_id = $1
AND s.status IN ('active', 'trialing', 'past_due')
AND s.current_period_start <= NOW()
AND s.current_period_end > NOW()
ORDER BY s.current_period_end DESC
LIMIT 1
"#,
)
.bind(user_id)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询订阅失败").with_source(err))?;
let (status, period_start, period_end, cancel_at_period_end, plan) = if let Some(sub) = sub {
(
sub.status,
sub.current_period_start,
sub.current_period_end,
sub.cancel_at_period_end,
SubscriptionPlanView {
id: sub.plan_id,
code: sub.plan_code,
name: sub.plan_name,
currency: sub.currency,
amount_cents: sub.amount_cents,
interval: sub.interval,
included_units_per_period: sub.included_units_per_period,
max_file_size_mb: sub.max_file_size_mb,
max_files_per_batch: sub.max_files_per_batch,
retention_days: sub.retention_days,
features: sub.features,
},
)
} else {
let plan: PlanView = sqlx::query_as::<_, PlanView>(
r#"
SELECT
id, code, name, currency, amount_cents, interval,
included_units_per_period, max_file_size_mb, max_files_per_batch, retention_days,
features
FROM plans
WHERE code = 'free'
LIMIT 1
"#,
)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "未找到 Free 套餐").with_source(err))?;
let (start, end) = billing::current_month_period_utc8(Utc::now());
(
"free".to_string(),
start,
end,
false,
SubscriptionPlanView {
id: plan.id,
code: plan.code,
name: plan.name,
currency: plan.currency,
amount_cents: plan.amount_cents,
interval: plan.interval,
included_units_per_period: plan.included_units_per_period,
max_file_size_mb: plan.max_file_size_mb,
max_files_per_batch: plan.max_files_per_batch,
retention_days: plan.retention_days,
features: plan.features,
},
)
};
Ok(Json(Envelope {
success: true,
data: SubscriptionResponse {
subscription: SubscriptionView {
status,
current_period_start: period_start,
current_period_end: period_end,
cancel_at_period_end,
plan,
},
},
}))
}
#[derive(Debug, Serialize)]
struct UsageResponse {
period_start: DateTime<Utc>,
period_end: DateTime<Utc>,
used_units: i64,
included_units: i64,
bonus_units: i64,
redeemed_units: i64,
total_units: i64,
remaining_units: i64,
}
async fn get_usage(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
) -> Result<Json<Envelope<UsageResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let billing = billing::get_user_billing(&state, user_id).await?;
let usage = quota::user_usage_balance(&state, &billing).await?;
Ok(Json(Envelope {
success: true,
data: UsageResponse {
period_start: billing.period_start,
period_end: billing.period_end,
used_units: usage.used_units,
included_units: usage.included_units,
bonus_units: usage.bonus_units,
redeemed_units: usage.redeemed_units,
total_units: usage.total_units,
remaining_units: usage.remaining_units,
},
}))
}
#[derive(Debug, Deserialize)]
struct PagingQuery {
page: Option<u32>,
limit: Option<u32>,
}
#[derive(Debug, FromRow, Serialize)]
struct InvoiceView {
invoice_number: String,
status: String,
currency: String,
total_amount_cents: i32,
period_start: Option<DateTime<Utc>>,
period_end: Option<DateTime<Utc>>,
hosted_invoice_url: Option<String>,
pdf_url: Option<String>,
paid_at: Option<DateTime<Utc>>,
created_at: DateTime<Utc>,
}
#[derive(Debug, Serialize)]
struct InvoicesResponse {
invoices: Vec<InvoiceView>,
page: u32,
limit: u32,
}
async fn list_invoices(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
axum::extract::Query(query): axum::extract::Query<PagingQuery>,
) -> Result<Json<Envelope<InvoicesResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let limit = query.limit.unwrap_or(20).clamp(1, 100);
let page = query.page.unwrap_or(1).max(1);
let offset = (page - 1) * limit;
let invoices = sqlx::query_as::<_, InvoiceView>(
r#"
SELECT
invoice_number,
status::text AS status,
currency,
total_amount_cents,
period_start,
period_end,
hosted_invoice_url,
pdf_url,
paid_at,
created_at
FROM invoices
WHERE user_id = $1
ORDER BY created_at DESC
LIMIT $2 OFFSET $3
"#,
)
.bind(user_id)
.bind(limit as i64)
.bind(offset as i64)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询发票失败").with_source(err))?;
Ok(Json(Envelope {
success: true,
data: InvoicesResponse {
invoices,
page,
limit,
},
}))
}
#[derive(Debug, Deserialize)]
struct CheckoutRequest {
plan_id: Uuid,
}
#[derive(Debug, Serialize, Deserialize)]
struct CheckoutResponse {
checkout_url: String,
}
async fn create_checkout(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
Json(req): Json<CheckoutRequest>,
) -> Result<Json<Envelope<CheckoutResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let idempotency_key = headers
.get("idempotency-key")
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|v| !v.is_empty())
.map(str::to_string);
let request_hash = idempotency_key.as_ref().map(|_| {
let plan = req.plan_id.to_string();
idempotency::sha256_hex(&[b"billing_checkout", plan.as_bytes()])
});
let mut idem_acquired = false;
if let (Some(idem), Some(request_hash)) = (idempotency_key.as_deref(), request_hash.as_deref())
{
match idempotency::begin(
&state,
idempotency::Scope::User(user_id),
idem,
request_hash,
state.config.idempotency_ttl_hours as i64,
)
.await?
{
idempotency::BeginResult::Replay { response_body, .. } => {
let resp: CheckoutResponse =
serde_json::from_value(response_body).map_err(|err| {
AppError::new(ErrorCode::Internal, "幂等结果解析失败").with_source(err)
})?;
return Ok(Json(Envelope {
success: true,
data: resp,
}));
}
idempotency::BeginResult::InProgress => {
if let Some((_status, body)) = idempotency::wait_for_replay(
&state,
idempotency::Scope::User(user_id),
idem,
request_hash,
10_000,
)
.await?
{
let resp: CheckoutResponse = serde_json::from_value(body).map_err(|err| {
AppError::new(ErrorCode::Internal, "幂等结果解析失败").with_source(err)
})?;
return Ok(Json(Envelope {
success: true,
data: resp,
}));
}
return Err(AppError::new(
ErrorCode::InvalidRequest,
"请求正在处理中,请稍后重试",
));
}
idempotency::BeginResult::Acquired => {
idem_acquired = true;
}
}
}
let session_result: Result<String, AppError> = (async {
let stripe_secret = settings::get_stripe_secret(&state)
.await
.map_err(|err| err.with_source("stripe secret not configured"))?;
#[derive(Debug, FromRow)]
struct PlanStripeRow {
stripe_price_id: Option<String>,
amount_cents: i32,
is_active: bool,
}
let plan = sqlx::query_as::<_, PlanStripeRow>(
r#"
SELECT stripe_price_id, amount_cents, is_active
FROM plans
WHERE id = $1
"#,
)
.bind(req.plan_id)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询套餐失败").with_source(err))?
.ok_or_else(|| AppError::new(ErrorCode::NotFound, "套餐不存在"))?;
if !plan.is_active {
return Err(AppError::new(ErrorCode::Forbidden, "套餐不可用"));
}
let Some(price_id) = plan.stripe_price_id.filter(|v| !v.trim().is_empty()) else {
return Err(AppError::new(ErrorCode::InvalidRequest, "该套餐不可订阅"));
};
if plan.amount_cents <= 0 {
return Err(AppError::new(ErrorCode::InvalidRequest, "该套餐不可订阅"));
}
#[derive(Debug, FromRow)]
struct UserStripeRow {
email: String,
billing_customer_id: Option<String>,
}
let user = sqlx::query_as::<_, UserStripeRow>(
"SELECT email, billing_customer_id FROM users WHERE id = $1",
)
.bind(user_id)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询用户失败").with_source(err))?;
let customer_id = if let Some(cus) = user.billing_customer_id {
cus
} else {
let cus = stripe_create_customer(&stripe_secret, &user.email, user_id).await?;
let _ = sqlx::query("UPDATE users SET billing_customer_id = $2 WHERE id = $1")
.bind(user_id)
.bind(&cus)
.execute(&state.db)
.await;
cus
};
let success_url = format!(
"{}/dashboard/billing?checkout=success",
state.config.public_base_url
);
let cancel_url = format!("{}/pricing?checkout=cancel", state.config.public_base_url);
stripe_create_checkout_session(
&stripe_secret,
&customer_id,
&price_id,
&success_url,
&cancel_url,
user_id,
)
.await
})
.await;
match session_result {
Ok(session) => {
if let (Some(idem), Some(request_hash)) =
(idempotency_key.as_deref(), request_hash.as_deref())
{
if idem_acquired {
let _ = idempotency::complete(
&state,
idempotency::Scope::User(user_id),
idem,
request_hash,
200,
serde_json::to_value(&CheckoutResponse {
checkout_url: session.clone(),
})
.unwrap_or(serde_json::Value::Null),
)
.await;
}
}
Ok(Json(Envelope {
success: true,
data: CheckoutResponse {
checkout_url: session,
},
}))
}
Err(err) => {
if let (Some(idem), Some(request_hash)) =
(idempotency_key.as_deref(), request_hash.as_deref())
{
if idem_acquired {
let _ = idempotency::abort(
&state,
idempotency::Scope::User(user_id),
idem,
request_hash,
)
.await;
}
}
Err(err)
}
}
}
#[derive(Debug, Serialize, Deserialize)]
struct PortalResponse {
url: String,
}
async fn create_portal(
State(state): State<AppState>,
jar: axum_extra::extract::cookie::CookieJar,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
headers: HeaderMap,
) -> Result<Json<Envelope<PortalResponse>>, AppError> {
let ip = context::client_ip(&headers, addr.ip());
let (_jar, principal) = context::authenticate(&state, jar, &headers, ip).await?;
let user_id = match principal {
context::Principal::User { user_id, .. } => user_id,
_ => return Err(AppError::new(ErrorCode::Unauthorized, "未登录")),
};
let stripe_secret = settings::get_stripe_secret(&state)
.await
.map_err(|err| err.with_source("stripe secret not configured"))?;
let customer_id: Option<String> =
sqlx::query_scalar("SELECT billing_customer_id FROM users WHERE id = $1")
.bind(user_id)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询用户失败").with_source(err))?;
let Some(customer_id) = customer_id.filter(|v| !v.trim().is_empty()) else {
return Err(AppError::new(
ErrorCode::InvalidRequest,
"未找到 Stripe Customer",
));
};
let return_url = format!("{}/dashboard/billing", state.config.public_base_url);
let url = stripe_create_portal_session(&stripe_secret, &customer_id, &return_url).await?;
Ok(Json(Envelope {
success: true,
data: PortalResponse { url },
}))
}
async fn stripe_create_customer(
secret: &str,
email: &str,
user_id: Uuid,
) -> Result<String, AppError> {
let resp: serde_json::Value = stripe_post_form(
secret,
"/v1/customers",
vec![
("email".to_string(), email.to_string()),
("metadata[user_id]".to_string(), user_id.to_string()),
],
)
.await?;
let id = resp
.get("id")
.and_then(|v| v.as_str())
.ok_or_else(|| AppError::new(ErrorCode::Internal, "Stripe customer 创建失败"))?;
Ok(id.to_string())
}
async fn stripe_create_checkout_session(
secret: &str,
customer_id: &str,
price_id: &str,
success_url: &str,
cancel_url: &str,
user_id: Uuid,
) -> Result<String, AppError> {
let resp: serde_json::Value = stripe_post_form(
secret,
"/v1/checkout/sessions",
vec![
("mode".to_string(), "subscription".to_string()),
("customer".to_string(), customer_id.to_string()),
("line_items[0][price]".to_string(), price_id.to_string()),
("line_items[0][quantity]".to_string(), "1".to_string()),
("success_url".to_string(), success_url.to_string()),
("cancel_url".to_string(), cancel_url.to_string()),
("allow_promotion_codes".to_string(), "true".to_string()),
("client_reference_id".to_string(), user_id.to_string()),
("metadata[user_id]".to_string(), user_id.to_string()),
],
)
.await?;
let url = resp
.get("url")
.and_then(|v| v.as_str())
.ok_or_else(|| AppError::new(ErrorCode::Internal, "Stripe checkout 创建失败"))?;
Ok(url.to_string())
}
async fn stripe_create_portal_session(
secret: &str,
customer_id: &str,
return_url: &str,
) -> Result<String, AppError> {
let resp: serde_json::Value = stripe_post_form(
secret,
"/v1/billing_portal/sessions",
vec![
("customer".to_string(), customer_id.to_string()),
("return_url".to_string(), return_url.to_string()),
],
)
.await?;
let url = resp
.get("url")
.and_then(|v| v.as_str())
.ok_or_else(|| AppError::new(ErrorCode::Internal, "Stripe portal 创建失败"))?;
Ok(url.to_string())
}
async fn stripe_post_form(
secret: &str,
path: &str,
form: Vec<(String, String)>,
) -> Result<serde_json::Value, AppError> {
let url = format!("https://api.stripe.com{path}");
let client = reqwest::Client::new();
let resp = client
.post(url)
.bearer_auth(secret)
.form(&form)
.send()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "Stripe 请求失败").with_source(err))?;
let status = resp.status();
let body = resp.text().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "Stripe 响应读取失败").with_source(err)
})?;
if !status.is_success() {
tracing::error!(status = %status, body = %body, "Stripe API error");
return Err(AppError::new(ErrorCode::Internal, "Stripe API 调用失败"));
}
serde_json::from_str(&body)
.map_err(|err| AppError::new(ErrorCode::Internal, "Stripe 响应解析失败").with_source(err))
}