mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 08:36:37 +08:00
压缩历史至 95081aff
This commit is contained in:
@@ -0,0 +1,46 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/logger"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
)
|
||||
|
||||
func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
|
||||
auditLog := loginRequiredAuditLog{
|
||||
UserID: user.ID,
|
||||
Username: user.Username,
|
||||
ClientIP: c.ClientIP(),
|
||||
Method: c.Request.Method,
|
||||
Path: c.Request.URL.Path,
|
||||
RequestURI: c.Request.RequestURI,
|
||||
UserAgent: c.Request.UserAgent(),
|
||||
Referer: c.Request.Referer(),
|
||||
}
|
||||
auditJSON, err := json.Marshal(auditLog)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
|
||||
logger.InfoF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
|
||||
} else {
|
||||
logger.InfoF(ctx, "[LoginRequiredAudit] %s", auditJSON)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
var (
|
||||
oauthConf *oauth2.Config
|
||||
oidcVerifier *oidc.IDTokenVerifier
|
||||
)
|
||||
|
||||
func init() {
|
||||
cfg := config.Config.OAuth2
|
||||
|
||||
if cfg.Issuer != "" {
|
||||
ctx := context.Background()
|
||||
provider, err := oidc.NewProvider(ctx, cfg.Issuer)
|
||||
if err != nil {
|
||||
log.Printf("[OAuth] 初始化 OIDC Provider 失败: %v,将仅使用 OAuth2", err)
|
||||
} else {
|
||||
oidcVerifier = provider.Verifier(&oidc.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
})
|
||||
log.Printf("[OAuth] OIDC Provider 初始化成功: %s", cfg.Issuer)
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化 OAuth2 配置
|
||||
scopes := []string{"profile", "email"}
|
||||
if oidcVerifier != nil {
|
||||
// 启用 OIDC 时添加 openid scope
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
|
||||
oauthConf = &oauth2.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
ClientSecret: cfg.ClientSecret,
|
||||
RedirectURL: cfg.RedirectURI,
|
||||
Scopes: scopes,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: cfg.AuthorizationEndpoint,
|
||||
TokenURL: cfg.TokenEndpoint,
|
||||
AuthStyle: oauth2.AuthStyleAutoDetect,
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
UserNameKey = "username"
|
||||
UserIDKey = "user_id"
|
||||
UserObjKey = "user_obj"
|
||||
)
|
||||
|
||||
const (
|
||||
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
||||
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
||||
)
|
||||
@@ -0,0 +1,23 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
const (
|
||||
InvalidState = "非法登录请求"
|
||||
IDTokenVerifyFailed = "ID Token 验证失败"
|
||||
NonceMismatch = "nonce 不匹配,可能存在重放攻击"
|
||||
)
|
||||
@@ -0,0 +1,152 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"github.com/linux-do/credit/internal/common"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/otel_trace"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func GetUserIDFromSession(s sessions.Session) uint64 {
|
||||
userID, ok := s.Get(UserIDKey).(uint64)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return userID
|
||||
}
|
||||
|
||||
func GetUserIDFromContext(c *gin.Context) uint64 {
|
||||
session := sessions.Default(c)
|
||||
return GetUserIDFromSession(session)
|
||||
}
|
||||
|
||||
// doOAuth 执行 OAuth2/OIDC 认证流程
|
||||
func doOAuth(ctx context.Context, code string, nonce string) (*model.User, error) {
|
||||
ctx, span := otel_trace.Start(ctx, "OAuth")
|
||||
defer span.End()
|
||||
|
||||
// 使用授权码换取 Token
|
||||
token, err := oauthConf.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var userInfo model.OAuthUserInfo
|
||||
|
||||
if oidcVerifier != nil {
|
||||
if rawIDToken, ok := token.Extra("id_token").(string); ok {
|
||||
idToken, verifyErr := oidcVerifier.Verify(ctx, rawIDToken)
|
||||
if verifyErr != nil {
|
||||
err := fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
if nonce != "" && idToken.Nonce != nonce {
|
||||
span.SetStatus(codes.Error, NonceMismatch)
|
||||
return nil, errors.New(NonceMismatch)
|
||||
}
|
||||
if claimsErr := idToken.Claims(&userInfo); claimsErr != nil {
|
||||
span.SetStatus(codes.Error, claimsErr.Error())
|
||||
return nil, claimsErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if userInfo.GetID() == 0 {
|
||||
client := oauthConf.Client(ctx, token)
|
||||
resp, httpErr := client.Get(config.Config.OAuth2.UserEndpoint)
|
||||
if httpErr != nil {
|
||||
span.SetStatus(codes.Error, httpErr.Error())
|
||||
return nil, httpErr
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
responseData, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
span.SetStatus(codes.Error, readErr.Error())
|
||||
return nil, readErr
|
||||
}
|
||||
if unmarshalErr := json.Unmarshal(responseData, &userInfo); unmarshalErr != nil {
|
||||
span.SetStatus(codes.Error, unmarshalErr.Error())
|
||||
return nil, unmarshalErr
|
||||
}
|
||||
}
|
||||
|
||||
if !userInfo.Active {
|
||||
err = errors.New(common.BannedAccount)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var user model.User
|
||||
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var holder model.User
|
||||
if conflictErr := tx.Where("username = ? AND id != ?", userInfo.Username, userInfo.GetID()).First(&holder).Error; conflictErr == nil {
|
||||
// 存在冲突 -> 将占用者改名并注销
|
||||
newParams := map[string]interface{}{
|
||||
"username": fmt.Sprintf("%s已注销: %s", holder.Username, uuid.NewString()),
|
||||
"is_active": false,
|
||||
}
|
||||
if updateErr := tx.Model(&holder).Updates(newParams).Error; updateErr != nil {
|
||||
return updateErr
|
||||
}
|
||||
}
|
||||
|
||||
// 根据 ID 处理当前用户的 更新 或 创建
|
||||
if queryErr := tx.Where("id = ?", userInfo.GetID()).First(&user).Error; queryErr == nil {
|
||||
// 用户已存在 -> 更新信息
|
||||
if activeErr := user.CheckActive(); activeErr != nil {
|
||||
return activeErr
|
||||
}
|
||||
user.UpdateFromOAuthInfo(&userInfo)
|
||||
if saveErr := tx.Save(&user).Error; saveErr != nil {
|
||||
return saveErr
|
||||
}
|
||||
} else if errors.Is(queryErr, gorm.ErrRecordNotFound) {
|
||||
// 用户不存在 -> 创建新用户
|
||||
user = model.User{}
|
||||
if createErr := user.CreateWithInitialCredit(tx, &userInfo); createErr != nil {
|
||||
return createErr
|
||||
}
|
||||
} else {
|
||||
return queryErr
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/common"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/otel_trace"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
type loginRequiredAuditLog struct {
|
||||
UserID uint64 `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
ClientIP string `json:"client_ip"`
|
||||
Method string `json:"method"`
|
||||
Path string `json:"path"`
|
||||
RequestURI string `json:"request_uri"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Referer string `json:"referer"`
|
||||
}
|
||||
|
||||
func LoginRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// init trace
|
||||
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired")
|
||||
defer span.End()
|
||||
|
||||
// load user
|
||||
userId := GetUserIDFromContext(c)
|
||||
if userId <= 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
|
||||
return
|
||||
}
|
||||
|
||||
// load user from db to make sure is active
|
||||
var user model.User
|
||||
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userId, true).First(&user)
|
||||
if tx.Error != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error_msg": tx.Error.Error(), "data": nil})
|
||||
return
|
||||
}
|
||||
|
||||
// log
|
||||
LogForAudit(ctx, &user, c)
|
||||
|
||||
// set user info
|
||||
util.SetToContext(c, UserObjKey, &user)
|
||||
|
||||
if risk, ok := checkOpenAPIUserRisk(ctx, user.ID); ok {
|
||||
if blocked := applyOpenAPIUserRisk(c, risk); blocked {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// next
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/logger"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
const (
|
||||
openAPIRiskCacheKeyFormat = "openapi_risk:user:%d"
|
||||
minOpenAPIRiskCacheTTL = time.Hour
|
||||
|
||||
riskLevelHeader = "X-Credit-Risk-Level"
|
||||
riskLabelsHeader = "X-Credit-Risk-Labels"
|
||||
riskItemsHeader = "X-Credit-Risks"
|
||||
exposeHeader = "Access-Control-Expose-Headers"
|
||||
|
||||
riskBlockedCode = "RISK_BLOCKED"
|
||||
riskBlockedMsg = "账号存在风险"
|
||||
)
|
||||
|
||||
type openAPIUserRiskItem struct {
|
||||
Label string `json:"label"`
|
||||
Value string `json:"value"`
|
||||
Desc string `json:"desc"`
|
||||
}
|
||||
|
||||
type openAPIUserRiskResponse struct {
|
||||
Risky bool `json:"risky"`
|
||||
RiskLevel string `json:"risk_level"`
|
||||
Risks []openAPIUserRiskItem `json:"risks"`
|
||||
}
|
||||
|
||||
type riskBlockDetails struct {
|
||||
RiskLevel string `json:"risk_level"`
|
||||
RiskLabels []string `json:"risk_labels"`
|
||||
Risks []openAPIUserRiskItem `json:"risks"`
|
||||
}
|
||||
|
||||
func checkOpenAPIUserRisk(ctx context.Context, userID uint64) (*openAPIUserRiskResponse, bool) {
|
||||
cfg := config.Config.OpenAPIRisk
|
||||
if !cfg.Enabled || strings.TrimSpace(cfg.BaseURL) == "" {
|
||||
return nil, false
|
||||
}
|
||||
if db.Redis == nil {
|
||||
logger.ErrorF(ctx, "[OpenAPIRisk] redis is not initialized, skip risk check")
|
||||
return nil, false
|
||||
}
|
||||
|
||||
cacheKey := fmt.Sprintf(openAPIRiskCacheKeyFormat, userID)
|
||||
var cached openAPIUserRiskResponse
|
||||
if err := db.GetJSON(ctx, cacheKey, &cached); err == nil {
|
||||
return &cached, true
|
||||
} else if err != nil && !errors.Is(err, redis.Nil) {
|
||||
logger.ErrorF(ctx, "[OpenAPIRisk] read cache failed, skip risk check: %v", err)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
risk, err := fetchOpenAPIUserRisk(ctx, userID)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[OpenAPIRisk] fetch user risk failed, skip risk check: %v", err)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
if err := db.SetJSON(ctx, cacheKey, risk, openAPIRiskCacheTTL()); err != nil {
|
||||
logger.ErrorF(ctx, "[OpenAPIRisk] write cache failed, skip risk check: %v", err)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
return risk, true
|
||||
}
|
||||
|
||||
func fetchOpenAPIUserRisk(ctx context.Context, userID uint64) (*openAPIUserRiskResponse, error) {
|
||||
cfg := config.Config.OpenAPIRisk
|
||||
endpoint := fmt.Sprintf(
|
||||
"%s/api/open/v1/risk/users/%d",
|
||||
strings.TrimRight(cfg.BaseURL, "/"),
|
||||
userID,
|
||||
)
|
||||
|
||||
headers := map[string]string{
|
||||
"Accept": "application/json",
|
||||
}
|
||||
if cfg.Username != "" || cfg.Password != "" {
|
||||
token := base64.StdEncoding.EncodeToString([]byte(cfg.Username + ":" + cfg.Password))
|
||||
headers["Authorization"] = "Basic " + token
|
||||
}
|
||||
|
||||
resp, err := util.Request(ctx, http.MethodGet, endpoint, nil, headers, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("unexpected status code: %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var risk openAPIUserRiskResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&risk); err != nil {
|
||||
return nil, fmt.Errorf("decode response failed: %w", err)
|
||||
}
|
||||
|
||||
return &risk, nil
|
||||
}
|
||||
|
||||
func openAPIRiskCacheTTL() time.Duration {
|
||||
ttl := time.Duration(config.Config.OpenAPIRisk.CacheTTLSeconds) * time.Second
|
||||
if ttl < minOpenAPIRiskCacheTTL {
|
||||
return minOpenAPIRiskCacheTTL
|
||||
}
|
||||
return ttl
|
||||
}
|
||||
|
||||
func applyOpenAPIUserRisk(c *gin.Context, risk *openAPIUserRiskResponse) bool {
|
||||
if risk == nil || !risk.Risky {
|
||||
return false
|
||||
}
|
||||
|
||||
labels := riskLabels(risk)
|
||||
items := riskItems(risk)
|
||||
cfg := config.Config.OpenAPIRisk
|
||||
if containsString(cfg.BlockRiskLevels, risk.RiskLevel) {
|
||||
setRiskHeaders(c, risk.RiskLevel, labels, items)
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
|
||||
"error_code": riskBlockedCode,
|
||||
"error_msg": riskBlockedMsg,
|
||||
"details": riskBlockDetails{
|
||||
RiskLevel: risk.RiskLevel,
|
||||
RiskLabels: labels,
|
||||
Risks: items,
|
||||
},
|
||||
})
|
||||
return true
|
||||
}
|
||||
|
||||
if containsString(cfg.PromptRiskLevels, risk.RiskLevel) {
|
||||
setRiskHeaders(c, risk.RiskLevel, labels, items)
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func setRiskHeaders(c *gin.Context, riskLevel string, labels []string, items []openAPIUserRiskItem) {
|
||||
labelsJSON, err := json.Marshal(labels)
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[OpenAPIRisk] marshal risk labels failed: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
itemsJSON, err := json.Marshal(items)
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[OpenAPIRisk] marshal risk items failed: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
c.Header(riskLevelHeader, riskLevel)
|
||||
c.Header(riskLabelsHeader, base64.StdEncoding.EncodeToString(labelsJSON))
|
||||
c.Header(riskItemsHeader, base64.StdEncoding.EncodeToString(itemsJSON))
|
||||
appendExposeHeaders(c, riskLevelHeader, riskLabelsHeader, riskItemsHeader)
|
||||
}
|
||||
|
||||
func appendExposeHeaders(c *gin.Context, names ...string) {
|
||||
existing := c.Writer.Header().Get(exposeHeader)
|
||||
exposed := make([]string, 0, len(names)+1)
|
||||
if existing != "" {
|
||||
exposed = append(exposed, strings.Split(existing, ",")...)
|
||||
}
|
||||
exposed = append(exposed, names...)
|
||||
|
||||
seen := make(map[string]struct{}, len(exposed))
|
||||
normalized := make([]string, 0, len(exposed))
|
||||
for _, header := range exposed {
|
||||
header = strings.TrimSpace(header)
|
||||
if header == "" {
|
||||
continue
|
||||
}
|
||||
key := strings.ToLower(header)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
normalized = append(normalized, header)
|
||||
}
|
||||
|
||||
c.Header(exposeHeader, strings.Join(normalized, ", "))
|
||||
}
|
||||
|
||||
func riskLabels(risk *openAPIUserRiskResponse) []string {
|
||||
labels := make([]string, 0, len(risk.Risks))
|
||||
for _, item := range risk.Risks {
|
||||
label := strings.TrimSpace(item.Label)
|
||||
if label == "" {
|
||||
continue
|
||||
}
|
||||
labels = append(labels, label)
|
||||
}
|
||||
return labels
|
||||
}
|
||||
|
||||
func riskItems(risk *openAPIUserRiskResponse) []openAPIUserRiskItem {
|
||||
items := make([]openAPIUserRiskItem, 0, len(risk.Risks))
|
||||
for _, item := range risk.Risks {
|
||||
item.Label = strings.TrimSpace(item.Label)
|
||||
item.Value = strings.TrimSpace(item.Value)
|
||||
item.Desc = strings.TrimSpace(item.Desc)
|
||||
if item.Label == "" {
|
||||
continue
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func containsString(values []string, target string) bool {
|
||||
target = strings.TrimSpace(target)
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/service"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
"github.com/shopspring/decimal"
|
||||
)
|
||||
|
||||
// GetLoginURL godoc
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Router /api/v1/oauth/login [get]
|
||||
func GetLoginURL(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
|
||||
// 生成 state
|
||||
state := uuid.NewString()
|
||||
cmd := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), state, OAuthStateCacheKeyExpiration)
|
||||
if cmd.Err() != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(cmd.Err().Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// 构造登录 URL
|
||||
var authURL string
|
||||
if oidcVerifier != nil {
|
||||
// OIDC 模式:state 同时用作 nonce
|
||||
authURL = oauthConf.AuthCodeURL(state, oidc.Nonce(state))
|
||||
} else {
|
||||
// 纯 OAuth2 模式
|
||||
authURL = oauthConf.AuthCodeURL(state)
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OK(authURL))
|
||||
}
|
||||
|
||||
type CallbackRequest struct {
|
||||
State string `json:"state"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
// Callback godoc
|
||||
// @Tags oauth
|
||||
// @Param request body CallbackRequest true "request body"
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Router /api/v1/oauth/callback [post]
|
||||
func Callback(c *gin.Context) {
|
||||
// 解析请求
|
||||
var req CallbackRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
|
||||
// 验证 state
|
||||
cmd := db.Redis.Get(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)))
|
||||
if cmd.Val() != req.State {
|
||||
c.JSON(http.StatusBadRequest, util.Err(InvalidState))
|
||||
return
|
||||
}
|
||||
db.Redis.Del(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)))
|
||||
|
||||
// 执行 OAuth/OIDC 认证
|
||||
user, err := doOAuth(ctx, req.Code, req.State)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
session := sessions.Default(c)
|
||||
session.Set(UserIDKey, user.ID)
|
||||
session.Set(UserNameKey, user.Username)
|
||||
if err := session.Save(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
LogForAudit(ctx, user, c)
|
||||
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
}
|
||||
|
||||
type BasicUserInfo struct {
|
||||
ID uint64 `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
TrustLevel model.TrustLevel `json:"trust_level"`
|
||||
AvatarUrl string `json:"avatar_url"`
|
||||
TotalReceive decimal.Decimal `json:"total_receive"`
|
||||
TotalPayment decimal.Decimal `json:"total_payment"`
|
||||
TotalTransfer decimal.Decimal `json:"total_transfer"`
|
||||
TotalCommunity decimal.Decimal `json:"total_community"`
|
||||
CommunityBalance decimal.Decimal `json:"community_balance"`
|
||||
AvailableBalance decimal.Decimal `json:"available_balance"`
|
||||
PendingBalance decimal.Decimal `json:"pending_balance"`
|
||||
PayScore int64 `json:"pay_score"`
|
||||
IsPayKey bool `json:"is_pay_key"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
RemainQuota decimal.Decimal `json:"remain_quota"`
|
||||
PayLevel model.PayLevel `json:"pay_level"`
|
||||
DailyLimit *int64 `json:"daily_limit"`
|
||||
}
|
||||
|
||||
// UserInfo godoc
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Router /api/v1/oauth/user-info [get]
|
||||
func UserInfo(c *gin.Context) {
|
||||
user, _ := util.GetFromContext[*model.User](c, UserObjKey)
|
||||
|
||||
var payConfig model.UserPayConfig
|
||||
if err := payConfig.GetByPayScore(db.DB(c.Request.Context()), user.PayScore); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// 计算剩余额度(-1 表示无限额)
|
||||
remainQuota := decimal.NewFromInt(-1)
|
||||
if payConfig.DailyLimit != nil && *payConfig.DailyLimit > 0 {
|
||||
todayUsed, err := service.GetTodayUsedAmount(db.DB(c.Request.Context()), user.ID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
remainQuota = decimal.NewFromInt(*payConfig.DailyLimit).Sub(todayUsed)
|
||||
}
|
||||
|
||||
c.JSON(
|
||||
http.StatusOK,
|
||||
util.OK(BasicUserInfo{
|
||||
ID: user.ID,
|
||||
Username: user.Username,
|
||||
Nickname: user.Nickname,
|
||||
TrustLevel: user.TrustLevel,
|
||||
AvatarUrl: user.AvatarUrl,
|
||||
TotalReceive: user.TotalReceive,
|
||||
TotalPayment: user.TotalPayment,
|
||||
TotalTransfer: user.TotalTransfer,
|
||||
TotalCommunity: user.TotalCommunity,
|
||||
CommunityBalance: user.CommunityBalance,
|
||||
AvailableBalance: user.AvailableBalance,
|
||||
PendingBalance: user.PendingBalance,
|
||||
PayScore: user.PayScore,
|
||||
IsPayKey: user.PayKey != "",
|
||||
IsAdmin: user.IsAdmin,
|
||||
RemainQuota: remainQuota,
|
||||
PayLevel: payConfig.Level,
|
||||
DailyLimit: payConfig.DailyLimit,
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
// Logout godoc
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Router /api/v1/oauth/logout [get]
|
||||
func Logout(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Options(util.GetSessionOptions(-1))
|
||||
session.Clear()
|
||||
if err := session.Save(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
}
|
||||
Reference in New Issue
Block a user