Files
gochat/backend/internal/service/enterprise_billing_worker.go
T
Rogee 851ca7e372 refactor: 移除 SAML/LDAP/MFA 登录方式,仅保留本地账号密码和 OIDC
后端移除:
- SAML: auth/saml.go, handler/saml_handler.go, account_saml_settings_handler.go,
  model/account_saml_settings.go, model/saml_idp_config.go, repo/*.go
- LDAP: auth/ldap.go, handler/ldap_handler.go, model/account_ldap_settings.go,
  repo/account_ldap_settings_repo.go
- MFA: auth/mfa.go, handler/mfa_handler.go
- auth_service: 移除 mfaService 依赖、MFARequired 字段、LoginWithMFA 方法
- auth_handler: 移除 LoginMFA handler、MFA 分支逻辑
- bootstrap: 移除 SAML/LDAP/MFA service 初始化和 handler 注册
- sso_middleware: 精简为仅支持 OIDC provider
- router: 移除 SAML/LDAP/MFA 路由注册
- config: 移除 SAMLConfig/LDAPConfig struct 和 defaults

前端移除:
- v3/login: 移除 MFA 验证流程和 SAML 登录入口
- v3/api/auth: 移除 MFA 响应处理
- v3/routes: 移除 SSO login 路由
- dashboard: 移除 MFA 设置页面、SAML 安全设置页面
- i18n: 移除 mfa.json
- featureFlags: 移除 SAML feature flag

.env.example / .env: 移除 SAML/LDAP 配置段
2026-07-29 19:03:04 +08:00

344 lines
11 KiB
Go

package service
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"sync"
"time"
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/worker"
"gorm.io/gorm"
)
const TaskTypeEnterpriseCreateStripeCustomer = "enterprise:create_stripe_customer"
type enterpriseCreateStripeCustomerJob struct {
AccountID uint `json:"account_id"`
}
var enterpriseBillingRegistrations sync.Map
func RegisterEnterpriseBillingJobs(wp *worker.WorkerPool, svc *AccountService) {
if wp == nil || svc == nil {
return
}
if _, loaded := enterpriseBillingRegistrations.LoadOrStore(wp, struct{}{}); loaded {
return
}
wp.Register(TaskTypeEnterpriseCreateStripeCustomer, svc.performCreateStripeCustomerJob)
}
func (s *AccountService) performCreateStripeCustomerJob(ctx context.Context, job *model.BackgroundJob) error {
var payload enterpriseCreateStripeCustomerJob
if err := json.Unmarshal(job.Payload, &payload); err != nil || payload.AccountID == 0 {
return fmt.Errorf("invalid Stripe customer job payload")
}
defer s.clearEnterpriseCustomerCreationFlag(ctx, payload.AccountID)
return s.createStripeCustomer(ctx, payload.AccountID)
}
func (s *AccountService) clearEnterpriseCustomerCreationFlag(ctx context.Context, accountID uint) {
db := s.repo.DB().WithContext(ctx)
if db.Dialector.Name() == "postgres" {
_ = db.Model(&model.Account{}).Where("id = ?", accountID).UpdateColumn("custom_attributes", gorm.Expr("custom_attributes - 'is_creating_customer'")).Error
return
}
var account model.Account
if db.First(&account, accountID).Error == nil {
attrs := account.CustomAttributesMap()
delete(attrs, "is_creating_customer")
if account.SetCustomAttributesMap(attrs) == nil {
_ = db.Model(&account).Update("custom_attributes", account.CustomAttributes).Error
}
}
}
func (s *AccountService) createStripeCustomer(ctx context.Context, accountID uint) error {
secret := strings.TrimSpace(os.Getenv("STRIPE_SECRET_KEY"))
if secret == "" {
return fmt.Errorf("STRIPE_SECRET_KEY is not configured")
}
account, err := s.repo.FindByID(ctx, accountID)
if err != nil {
return err
}
plan, err := s.defaultCloudPlan(ctx)
if err != nil {
return err
}
currency := s.accountBillingCurrency(ctx, account)
priceID := cloudPlanPriceID(plan, currency)
if priceID == "" {
return fmt.Errorf("default cloud plan has no Stripe price")
}
attrs := account.CustomAttributesMap()
customerID := ""
if value := attrs["stripe_customer_id"]; value != nil {
customerID = strings.TrimSpace(fmt.Sprint(value))
}
var subscription map[string]any
if customerID != "" {
subscription, err = s.stripeActiveSubscription(ctx, secret, customerID)
if err != nil {
return err
}
if subscription != nil && !cloudPlanContainsProduct(plan, stripeSubscriptionFields(subscription).productID) {
return nil
}
}
if customerID == "" {
customerID, err = s.stripeCreateCustomer(ctx, secret, account, currency)
if err != nil {
return err
}
}
if subscription == nil {
subscription, err = s.stripeCreateSubscription(ctx, secret, customerID, priceID, cloudPlanDefaultQuantity(plan))
if err != nil {
return err
}
}
fields := stripeSubscriptionFields(subscription)
planName := strings.TrimSpace(fmt.Sprint(plan["name"]))
attrs["stripe_customer_id"] = customerID
attrs["stripe_price_id"] = fields.priceID
attrs["stripe_product_id"] = fields.productID
attrs["plan_name"] = planName
attrs["subscribed_quantity"] = fields.quantity
attrs["subscription_status"] = fields.status
attrs["billing_currency"] = supportedBillingCurrency(fields.currency)
if fields.periodEnd > 0 {
attrs["subscription_ends_on"] = time.Unix(fields.periodEnd, 0).UTC().Format(time.RFC3339)
}
delete(attrs, "is_creating_customer")
if err := account.SetCustomAttributesMap(attrs); err != nil {
return err
}
reconcileCloudPlanFeatures(account, planName, planName)
return s.repo.Update(ctx, account)
}
func (s *AccountService) defaultCloudPlan(ctx context.Context) (map[string]any, error) {
var config model.InstallationConfig
if err := s.repo.DB().WithContext(ctx).Where("name = ?", "CHATWOOT_CLOUD_PLANS").First(&config).Error; err != nil {
return nil, err
}
var plans []map[string]any
if err := json.Unmarshal([]byte(config.Value), &plans); err != nil || len(plans) == 0 {
return nil, fmt.Errorf("CHATWOOT_CLOUD_PLANS is empty")
}
return plans[0], nil
}
func cloudPlanPriceID(plan map[string]any, currency string) string {
raw := plan["price_ids"]
if values, ok := raw.([]any); ok {
return firstString(values)
}
byCurrency, _ := raw.(map[string]any)
for _, key := range []string{currency, "usd"} {
if value := firstStringValue(byCurrency[key]); value != "" {
return value
}
}
for _, value := range byCurrency {
if priceID := firstStringValue(value); priceID != "" {
return priceID
}
}
return ""
}
func firstStringValue(value any) string {
if values, ok := value.([]any); ok {
return firstString(values)
}
return strings.TrimSpace(fmt.Sprint(value))
}
func firstString(values []any) string {
for _, value := range values {
if text := strings.TrimSpace(fmt.Sprint(value)); text != "" {
return text
}
}
return ""
}
func cloudPlanDefaultQuantity(plan map[string]any) int {
if quantity := int(numberValue(plan["default_quantity"])); quantity > 0 {
return quantity
}
return 2
}
func cloudPlanContainsProduct(plan map[string]any, productID string) bool {
if values, ok := plan["product_id"].([]any); ok {
for _, value := range values {
if fmt.Sprint(value) == productID {
return true
}
}
return false
}
return strings.TrimSpace(fmt.Sprint(plan["product_id"])) == productID
}
func (s *AccountService) stripeCreateCustomer(ctx context.Context, secret string, account *model.Account, currency string) (string, error) {
values := url.Values{"name": {account.Name}}
var admin model.User
_ = s.repo.DB().WithContext(ctx).Joins("JOIN account_users ON account_users.user_id = users.id").Where("account_users.account_id = ? AND account_users.role = ?", account.ID, "administrator").Order("account_users.id ASC").First(&admin).Error
values.Set("email", admin.Email)
if currency == "brl" {
values.Set("address[country]", "BR")
values.Set("preferred_locales[0]", "pt-BR")
}
response, err := stripeFormRequest(ctx, secret, http.MethodPost, "/v1/customers", values)
if err != nil {
return "", err
}
id := strings.TrimSpace(fmt.Sprint(response["id"]))
if id == "" {
return "", fmt.Errorf("Stripe customer response has no id")
}
return id, nil
}
func (s *AccountService) stripeActiveSubscription(ctx context.Context, secret, customerID string) (map[string]any, error) {
response, err := stripeFormRequest(ctx, secret, http.MethodGet, "/v1/subscriptions", url.Values{"customer": {customerID}, "status": {"active"}, "limit": {"1"}})
if err != nil {
return nil, err
}
data, _ := response["data"].([]any)
if len(data) == 0 {
return nil, nil
}
subscription, _ := data[0].(map[string]any)
return subscription, nil
}
func (s *AccountService) stripeCreateSubscription(ctx context.Context, secret, customerID, priceID string, quantity int) (map[string]any, error) {
return stripeFormRequest(ctx, secret, http.MethodPost, "/v1/subscriptions", url.Values{"customer": {customerID}, "items[0][price]": {priceID}, "items[0][quantity]": {strconv.Itoa(quantity)}})
}
func stripeFormRequest(ctx context.Context, secret, method, path string, values url.Values) (map[string]any, error) {
base := strings.TrimRight(os.Getenv("STRIPE_API_BASE"), "/")
if base == "" {
base = "https://api.stripe.com"
}
endpoint := base + path
var body *strings.Reader
if method == http.MethodGet {
endpoint += "?" + values.Encode()
body = strings.NewReader("")
} else {
body = strings.NewReader(values.Encode())
}
request, err := http.NewRequestWithContext(ctx, method, endpoint, body)
if err != nil {
return nil, err
}
request.SetBasicAuth(secret, "")
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
response, err := http.DefaultClient.Do(request)
if err != nil {
return nil, err
}
defer response.Body.Close()
var payload map[string]any
if err := json.NewDecoder(response.Body).Decode(&payload); err != nil {
return nil, err
}
if response.StatusCode < 200 || response.StatusCode >= 300 {
return nil, fmt.Errorf("Stripe API returned %d: %v", response.StatusCode, payload)
}
return payload, nil
}
type stripeSubscriptionData struct {
priceID, productID, currency, status string
quantity int
periodEnd int64
}
func stripeSubscriptionFields(subscription map[string]any) stripeSubscriptionData {
result := stripeSubscriptionData{status: strings.TrimSpace(fmt.Sprint(subscription["status"])), quantity: int(numberValue(subscription["quantity"])), periodEnd: int64(numberValue(subscription["current_period_end"]))}
price, _ := subscription["plan"].(map[string]any)
if items, ok := subscription["items"].(map[string]any); ok {
if data, ok := items["data"].([]any); ok && len(data) > 0 {
item, _ := data[0].(map[string]any)
if result.quantity == 0 {
result.quantity = int(numberValue(item["quantity"]))
}
if result.periodEnd == 0 {
result.periodEnd = int64(numberValue(item["current_period_end"]))
}
if current, ok := item["price"].(map[string]any); ok {
price = current
}
}
}
result.priceID = strings.TrimSpace(fmt.Sprint(price["id"]))
result.productID = strings.TrimSpace(fmt.Sprint(price["product"]))
result.currency = strings.TrimSpace(fmt.Sprint(price["currency"]))
return result
}
func numberValue(value any) float64 {
switch number := value.(type) {
case float64:
return number
case int:
return float64(number)
case int64:
return float64(number)
default:
parsed, _ := strconv.ParseFloat(fmt.Sprint(value), 64)
return parsed
}
}
func supportedBillingCurrency(currency string) string {
currency = strings.ToLower(strings.TrimSpace(currency))
if _, ok := supportedBillingCurrencies[currency]; ok {
return currency
}
return "usd"
}
func reconcileCloudPlanFeatures(account *model.Account, planName, defaultPlanName string) {
flags := map[string]bool{}
_ = json.Unmarshal([]byte(account.FeatureFlags), &flags)
startup := []string{"inbound_emails", "help_center", "campaigns", "team_management", "channel_facebook", "channel_email", "channel_instagram", "channel_tiktok", "captain_integration", "captain_document_auto_sync", "advanced_search_indexing", "advanced_search", "linear_integration", "channel_voice"}
business := []string{"sla", "custom_roles", "csat_review_notes", "conversation_required_attributes", "advanced_assignment", "custom_tools", "companies"}
enterprise := []string{"audit_logs", "disable_branding", "oidc"}
for _, feature := range append(append(append([]string{}, startup...), business...), enterprise...) {
flags[feature] = false
}
flags["captain_integration_v2"] = false
if planName != defaultPlanName {
for _, feature := range startup {
flags[feature] = true
}
if planName == "Business" || planName == "Enterprise" {
for _, feature := range business {
flags[feature] = true
}
}
if planName == "Enterprise" {
for _, feature := range enterprise {
flags[feature] = true
}
}
}
encoded, _ := json.Marshal(flags)
account.FeatureFlags = string(encoded)
}