mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 15:06:37 +08:00
oauth
This commit is contained in:
@@ -0,0 +1,215 @@
|
||||
package auth_source
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
type AuthSourceRequest struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
IsActive bool `json:"is_active"`
|
||||
ClientID string `json:"client_id"`
|
||||
ClientSecret string `json:"client_secret"`
|
||||
OpenIDDiscoveryURL string `json:"openid_discovery_url"`
|
||||
Scopes string `json:"scopes"`
|
||||
IconURL string `json:"icon_url"`
|
||||
}
|
||||
|
||||
type ToggleAuthSourceRequest struct {
|
||||
IsActive bool `json:"is_active"`
|
||||
}
|
||||
|
||||
// ListAuthSources 获取认证源列表
|
||||
// @Summary 获取认证源列表
|
||||
// @Description 返回所有已配置的 OAuth/OIDC 认证源列表,包括已启用和未启用的,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=[]model.AuthSource} "认证源列表"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/auth-sources [get]
|
||||
func ListAuthSources(c *gin.Context) {
|
||||
sources, err := model.GetAuthSources()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OK(sources))
|
||||
}
|
||||
|
||||
// CreateAuthSource 创建认证源
|
||||
// @Summary 创建认证源
|
||||
// @Description 创建一个新的 OAuth/OIDC 认证源配置,认证源名称必须唯一且符合命名规范,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body auth_source.AuthSourceRequest true "创建认证源参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=model.AuthSource} "创建成功,返回认证源信息"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误或验证失败"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/auth-sources [post]
|
||||
func CreateAuthSource(c *gin.Context) {
|
||||
var req AuthSourceRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
source := model.AuthSource{
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
DisplayName: req.DisplayName,
|
||||
IsActive: req.IsActive,
|
||||
ClientID: req.ClientID,
|
||||
ClientSecret: req.ClientSecret,
|
||||
OpenIDDiscoveryURL: req.OpenIDDiscoveryURL,
|
||||
Scopes: req.Scopes,
|
||||
IconURL: req.IconURL,
|
||||
}
|
||||
if err := model.CreateAuthSource(&source); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
source.Sanitize()
|
||||
c.JSON(http.StatusOK, util.OK(source))
|
||||
}
|
||||
|
||||
// UpdateAuthSource 更新认证源
|
||||
// @Summary 更新认证源
|
||||
// @Description 更新指定 ID 的认证源配置。若 client_secret 字段为空,则保留原有密钥不变,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "认证源 ID 或名称"
|
||||
// @Param request body auth_source.AuthSourceRequest true "更新认证源参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=model.AuthSource} "更新成功,返回更新后的认证源信息"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误或验证失败"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/auth-sources/{id} [put]
|
||||
func UpdateAuthSource(c *gin.Context) {
|
||||
id, err := parseSourceID(c)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
var req AuthSourceRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
source := model.AuthSource{
|
||||
ID: id,
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
DisplayName: req.DisplayName,
|
||||
IsActive: req.IsActive,
|
||||
ClientID: req.ClientID,
|
||||
ClientSecret: req.ClientSecret,
|
||||
OpenIDDiscoveryURL: req.OpenIDDiscoveryURL,
|
||||
Scopes: req.Scopes,
|
||||
IconURL: req.IconURL,
|
||||
}
|
||||
keepSecret := source.ClientSecret == ""
|
||||
if err := model.UpdateAuthSource(&source, keepSecret); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
updated, err := model.GetAuthSourceByID(id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
updated.Sanitize()
|
||||
c.JSON(http.StatusOK, util.OK(updated))
|
||||
}
|
||||
|
||||
// ToggleAuthSource 切换认证源启用状态
|
||||
// @Summary 切换认证源启用状态
|
||||
// @Description 启用或禁用指定认证源。尝试启用时将验证 Client ID 和 Client Secret 是否已配置,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "认证源 ID 或名称"
|
||||
// @Param request body auth_source.ToggleAuthSourceRequest true "启用状态"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "切换成功"
|
||||
// @Failure 400 {object} util.ResponseAny "验证失败或认证源不存在"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/auth-sources/{id}/toggle [put]
|
||||
func ToggleAuthSource(c *gin.Context) {
|
||||
id, err := parseSourceID(c)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
var req ToggleAuthSourceRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := model.ToggleAuthSource(id, req.IsActive); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
}
|
||||
|
||||
// DeleteAuthSource 删除认证源
|
||||
// @Summary 删除认证源
|
||||
// @Description 删除指定认证源及其关联的所有外部帐号绑定记录,警告:删除后相关用户将无法通过该源登录,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "认证源 ID 或名称"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "删除成功"
|
||||
// @Failure 400 {object} util.ResponseAny "ID 无效或删除失败"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/auth-sources/{id} [delete]
|
||||
func DeleteAuthSource(c *gin.Context) {
|
||||
id, err := parseSourceID(c)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := model.DeleteAuthSource(id); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
}
|
||||
|
||||
func parseSourceID(c *gin.Context) (uint64, error) {
|
||||
raw := c.Param("id")
|
||||
if raw == "" {
|
||||
return 0, errors.New("认证源 ID 无效")
|
||||
}
|
||||
source, err := model.GetAuthSourceByName(raw)
|
||||
if err == nil {
|
||||
return source.ID, nil
|
||||
}
|
||||
var id uint64
|
||||
if _, scanErr := fmt.Sscanf(raw, "%d", &id); scanErr != nil || id == 0 {
|
||||
return 0, errors.New("认证源 ID 无效")
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
/*
|
||||
Copyright 2026 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 auth_source
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/auth-sources", ListAuthSources)
|
||||
adminGroup.POST("/auth-sources", CreateAuthSource)
|
||||
adminGroup.PUT("/auth-sources/:id", UpdateAuthSource)
|
||||
adminGroup.PUT("/auth-sources/:id/toggle", ToggleAuthSource)
|
||||
adminGroup.DELETE("/auth-sources/:id", DeleteAuthSource)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestListAuthSources(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed source
|
||||
source := model.AuthSource{
|
||||
ID: 1,
|
||||
Name: "google",
|
||||
Type: "oidc",
|
||||
DisplayName: "Google Auth",
|
||||
IsActive: true,
|
||||
ClientID: "client_id_123",
|
||||
ClientSecret: "client_secret_456",
|
||||
OpenIDDiscoveryURL: "https://accounts.google.com",
|
||||
}
|
||||
dbConn.Create(&source)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/auth-sources", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var sources []model.AuthSource
|
||||
json.Unmarshal(dataBytes, &sources)
|
||||
|
||||
if len(sources) != 1 {
|
||||
t.Errorf("expected 1 auth source, got %d", len(sources))
|
||||
}
|
||||
if sources[0].Name != "google" {
|
||||
t.Errorf("expected name 'google', got '%s'", sources[0].Name)
|
||||
}
|
||||
// Verify sanitize removed the secret
|
||||
if sources[0].ClientSecret != "" {
|
||||
t.Error("client secret should be sanitized")
|
||||
}
|
||||
if !sources[0].ClientSecretConfigured {
|
||||
t.Error("client secret configured flag should be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateAuthSource(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("create successfully", func(t *testing.T) {
|
||||
reqPayload := AuthSourceRequest{
|
||||
Name: "github",
|
||||
Type: "oidc",
|
||||
DisplayName: "GitHub OIDC",
|
||||
IsActive: true,
|
||||
ClientID: "client_id_gh",
|
||||
ClientSecret: "client_secret_gh",
|
||||
OpenIDDiscoveryURL: "https://github.com",
|
||||
}
|
||||
body, _ := json.Marshal(reqPayload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// Verify database
|
||||
var src model.AuthSource
|
||||
dbConn.Where("name = ?", "github").First(&src)
|
||||
if src.ClientID != "client_id_gh" {
|
||||
t.Errorf("expected client_id_gh, got '%s'", src.ClientID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("create invalid validation failure", func(t *testing.T) {
|
||||
reqPayload := AuthSourceRequest{
|
||||
Name: "invalid name!",
|
||||
Type: "oidc",
|
||||
DisplayName: "Invalid",
|
||||
IsActive: true,
|
||||
ClientID: "client_id_val",
|
||||
ClientSecret: "client_secret_val",
|
||||
OpenIDDiscoveryURL: "https://discovery.url",
|
||||
}
|
||||
body, _ := json.Marshal(reqPayload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateAuthSource(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed source
|
||||
source := model.AuthSource{
|
||||
ID: 1,
|
||||
Name: "microsoft",
|
||||
Type: "oidc",
|
||||
DisplayName: "Microsoft",
|
||||
IsActive: true,
|
||||
ClientID: "old_client_id",
|
||||
ClientSecret: "old_secret",
|
||||
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
|
||||
}
|
||||
dbConn.Create(&source)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("update keep client secret", func(t *testing.T) {
|
||||
reqPayload := AuthSourceRequest{
|
||||
Name: "microsoft",
|
||||
Type: "oidc",
|
||||
DisplayName: "Microsoft Updated",
|
||||
IsActive: true,
|
||||
ClientID: "new_client_id",
|
||||
ClientSecret: "", // empty implies keeping existing secret
|
||||
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
|
||||
}
|
||||
body, _ := json.Marshal(reqPayload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var src model.AuthSource
|
||||
dbConn.First(&src, 1)
|
||||
if src.DisplayName != "Microsoft Updated" {
|
||||
t.Errorf("expected display name update, got '%s'", src.DisplayName)
|
||||
}
|
||||
if src.ClientSecret != "old_secret" {
|
||||
t.Errorf("expected old secret to be preserved, got '%s'", src.ClientSecret)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("update new client secret", func(t *testing.T) {
|
||||
reqPayload := AuthSourceRequest{
|
||||
Name: "microsoft",
|
||||
Type: "oidc",
|
||||
DisplayName: "Microsoft Updated Again",
|
||||
IsActive: true,
|
||||
ClientID: "new_client_id",
|
||||
ClientSecret: "brand_new_secret",
|
||||
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
|
||||
}
|
||||
body, _ := json.Marshal(reqPayload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/microsoft", bytes.NewBuffer(body)) // Using Name instead of ID
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var src model.AuthSource
|
||||
dbConn.First(&src, 1)
|
||||
if src.ClientSecret != "brand_new_secret" {
|
||||
t.Errorf("expected secret update, got '%s'", src.ClientSecret)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestToggleAuthSource(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
source := model.AuthSource{
|
||||
ID: 1,
|
||||
Name: "test_source",
|
||||
Type: "oidc",
|
||||
DisplayName: "Test Source",
|
||||
IsActive: false,
|
||||
ClientID: "",
|
||||
ClientSecret: "",
|
||||
OpenIDDiscoveryURL: "https://test.discovery.url",
|
||||
}
|
||||
dbConn.Create(&source)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("cannot activate without credentials", func(t *testing.T) {
|
||||
payload := ToggleAuthSourceRequest{IsActive: true}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request when activating without client_id/secret, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("toggle success after setting credentials", func(t *testing.T) {
|
||||
// Set credentials first
|
||||
dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Updates(map[string]interface{}{
|
||||
"client_id": "id",
|
||||
"client_secret": "secret",
|
||||
})
|
||||
|
||||
payload := ToggleAuthSourceRequest{IsActive: true}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var src model.AuthSource
|
||||
dbConn.First(&src, 1)
|
||||
if !src.IsActive {
|
||||
t.Error("auth source should be activated")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteAuthSource(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
source := model.AuthSource{
|
||||
ID: 1,
|
||||
Name: "delete_me",
|
||||
Type: "oidc",
|
||||
DisplayName: "Delete Me",
|
||||
IsActive: true,
|
||||
ClientID: "id",
|
||||
ClientSecret: "secret",
|
||||
OpenIDDiscoveryURL: "https://delete.me",
|
||||
}
|
||||
dbConn.Create(&source)
|
||||
|
||||
externalAccount := model.ExternalAccount{
|
||||
ID: 10,
|
||||
AuthSourceID: 1,
|
||||
UserID: 50,
|
||||
ExternalID: "ext_50",
|
||||
}
|
||||
dbConn.Create(&externalAccount)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/auth-sources/1", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
// Verify AuthSource is deleted
|
||||
var srcCount int64
|
||||
dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Count(&srcCount)
|
||||
if srcCount != 0 {
|
||||
t.Error("AuthSource should be deleted from the database")
|
||||
}
|
||||
|
||||
// Verify ExternalAccount bindings are also deleted
|
||||
var extCount int64
|
||||
dbConn.Model(&model.ExternalAccount{}).Where("auth_source_id = ?", 1).Count(&extCount)
|
||||
if extCount != 0 {
|
||||
t.Error("related ExternalAccount bindings should be deleted")
|
||||
}
|
||||
}
|
||||
@@ -47,8 +47,13 @@ type UpdateSystemConfigRequest struct {
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body CreateSystemConfigRequest true "创建请求参数"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Security SessionCookie
|
||||
// @Param request body system_config.CreateSystemConfigRequest true "创建请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "创建成功"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误或配置键已存在"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/system-configs [post]
|
||||
func CreateSystemConfig(c *gin.Context) {
|
||||
var req CreateSystemConfigRequest
|
||||
@@ -98,8 +103,12 @@ func CreateSystemConfig(c *gin.Context) {
|
||||
// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param type query string false "配置类型(system/business)"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=[]model.SystemConfig} "系统配置列表"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/system-configs [get]
|
||||
func ListSystemConfigs(c *gin.Context) {
|
||||
configType := c.Query("type")
|
||||
@@ -122,8 +131,13 @@ func ListSystemConfigs(c *gin.Context) {
|
||||
// @Description 根据配置键获取对应的系统配置详情,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "配置键"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=model.SystemConfig} "系统配置详情"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 404 {object} util.ResponseAny "配置不存在"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/system-configs/{key} [get]
|
||||
func GetSystemConfig(c *gin.Context) {
|
||||
var config model.SystemConfig
|
||||
@@ -145,9 +159,15 @@ func GetSystemConfig(c *gin.Context) {
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "配置键"
|
||||
// @Param request body UpdateSystemConfigRequest true "更新请求参数"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Param request body system_config.UpdateSystemConfigRequest true "更新请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "更新成功"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 404 {object} util.ResponseAny "配置不存在"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/system-configs/{key} [put]
|
||||
func UpdateSystemConfig(c *gin.Context) {
|
||||
var req UpdateSystemConfigRequest
|
||||
@@ -197,8 +217,13 @@ func UpdateSystemConfig(c *gin.Context) {
|
||||
// @Description 根据配置键删除对应配置,同时从 Redis 中移除对应缓存,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "配置键"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "删除成功"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 404 {object} util.ResponseAny "配置不存在"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/system-configs/{key} [delete]
|
||||
func DeleteSystemConfig(c *gin.Context) {
|
||||
key := c.Param("key")
|
||||
|
||||
@@ -0,0 +1,302 @@
|
||||
/*
|
||||
Copyright 2026 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 system_config
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.POST("/system-configs", CreateSystemConfig)
|
||||
adminGroup.GET("/system-configs", ListSystemConfigs)
|
||||
|
||||
systemConfigRouter := adminGroup.Group("/system-configs/:key")
|
||||
{
|
||||
systemConfigRouter.GET("", GetSystemConfig)
|
||||
systemConfigRouter.PUT("", UpdateSystemConfig)
|
||||
systemConfigRouter.DELETE("", DeleteSystemConfig)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func TestCreateSystemConfig(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("create successfully", func(t *testing.T) {
|
||||
payload := CreateSystemConfigRequest{
|
||||
Key: "custom_key",
|
||||
Value: "custom_value",
|
||||
Type: "system",
|
||||
Description: "desc",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// Verify database
|
||||
var cfg model.SystemConfig
|
||||
err := dbConn.Where("key = ?", "custom_key").First(&cfg).Error
|
||||
if err != nil {
|
||||
t.Fatalf("failed to find system config in DB: %v", err)
|
||||
}
|
||||
|
||||
// Verify Redis Cache
|
||||
var redisConfig model.SystemConfig
|
||||
err = db.HGetJSON(context.Background(), model.SystemConfigRedisHashKey, "custom_key", &redisConfig)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to find system config in Redis: %v", err)
|
||||
}
|
||||
if redisConfig.Value != "custom_value" {
|
||||
t.Errorf("expected value 'custom_value', got '%s'", redisConfig.Value)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("create duplicate key error", func(t *testing.T) {
|
||||
// Key "custom_key" already exists from previous test
|
||||
payload := CreateSystemConfigRequest{
|
||||
Key: "custom_key",
|
||||
Value: "another_value",
|
||||
Type: "system",
|
||||
Description: "desc",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestListSystemConfigs(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("list all seeded configurations", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var configs []model.SystemConfig
|
||||
json.Unmarshal(dataBytes, &configs)
|
||||
|
||||
// Defaults seed 7 configurations
|
||||
if len(configs) != 7 {
|
||||
t.Errorf("expected 7 default configs, got %d", len(configs))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by type business", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs?type=business", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var configs []model.SystemConfig
|
||||
json.Unmarshal(dataBytes, &configs)
|
||||
|
||||
if len(configs) != 1 || configs[0].Key != model.ConfigKeyMaxAPIKeysPerUser {
|
||||
t.Errorf("expected 1 business config (max_api_keys_per_user), got %d: %v", len(configs), configs)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetSystemConfig(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("get existing configuration", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var cfg model.SystemConfig
|
||||
json.Unmarshal(dataBytes, &cfg)
|
||||
|
||||
if cfg.Value != "Antigravity Project" {
|
||||
t.Errorf("expected 'Antigravity Project', got '%s'", cfg.Value)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("get non-existent config", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/non_existent_key", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateSystemConfig(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("update successfully", func(t *testing.T) {
|
||||
payload := UpdateSystemConfigRequest{
|
||||
Value: "Super Site Name",
|
||||
Description: "Updated Description",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// Verify database
|
||||
var cfg model.SystemConfig
|
||||
dbConn.Where("key = ?", model.ConfigKeySiteName).First(&cfg)
|
||||
if cfg.Value != "Super Site Name" || cfg.Description != "Updated Description" {
|
||||
t.Errorf("database values not updated: %+v", cfg)
|
||||
}
|
||||
|
||||
// Verify Redis
|
||||
var redisConfig model.SystemConfig
|
||||
db.HGetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeySiteName, &redisConfig)
|
||||
if redisConfig.Value != "Super Site Name" {
|
||||
t.Errorf("redis cache value not updated, got '%s'", redisConfig.Value)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("update non-existent config", func(t *testing.T) {
|
||||
payload := UpdateSystemConfigRequest{
|
||||
Value: "New Value",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/invalid_key", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteSystemConfig(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("delete successfully", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// Verify database
|
||||
var count int64
|
||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeySiteName).Count(&count)
|
||||
if count != 0 {
|
||||
t.Error("config still exists in DB")
|
||||
}
|
||||
|
||||
// Verify Redis Cache removal
|
||||
var redisConfig model.SystemConfig
|
||||
err := db.HGetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeySiteName, &redisConfig)
|
||||
if err == nil {
|
||||
t.Error("config should have been deleted from Redis cache")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delete non-existent config", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/system-configs/invalid_key", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -30,10 +30,13 @@ import (
|
||||
|
||||
// ListTaskTypes 获取支持的任务类型列表
|
||||
// @Summary 获取支持的任务类型
|
||||
// @Description 返回系统支持的所有可调度任务类型列表,需要管理员权限
|
||||
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=[]task.TaskMeta} "任务类型列表"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/types [get]
|
||||
func ListTaskTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, util.OK(task.DispatchableTasks))
|
||||
@@ -53,8 +56,13 @@ type DispatchTaskRequest struct {
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body DispatchTaskRequest true "任务请求参数"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Security SessionCookie
|
||||
// @Param request body task.DispatchTaskRequest true "任务请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "任务已入队"
|
||||
// @Failure 400 {object} util.ResponseAny "任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "任务入队失败"
|
||||
// @Router /api/v1/admin/tasks/dispatch [post]
|
||||
func DispatchTask(c *gin.Context) {
|
||||
var req DispatchTaskRequest
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
/*
|
||||
Copyright 2026 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 task
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/task"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/tasks/types", ListTaskTypes)
|
||||
adminGroup.POST("/tasks/dispatch", DispatchTask)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestListTaskTypes(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/types", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var taskMetas []task.TaskMeta
|
||||
json.Unmarshal(dataBytes, &taskMetas)
|
||||
|
||||
if len(taskMetas) == 0 {
|
||||
t.Error("expected at least one dispatchable task type")
|
||||
}
|
||||
|
||||
foundCleanup := false
|
||||
for _, m := range taskMetas {
|
||||
if m.Type == task.TaskTypeCleanupUploads {
|
||||
foundCleanup = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundCleanup {
|
||||
t.Errorf("expected task type %s to be listed", task.TaskTypeCleanupUploads)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDispatchTask(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true, SignKey: "admin_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("dispatch valid task successfully", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: task.TaskTypeCleanupUploads,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("dispatch invalid task type failure", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: "invalid_task_type",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
if resp.ErrorMsg != InvalidTaskType {
|
||||
t.Errorf("expected error message '%s', got '%s'", InvalidTaskType, resp.ErrorMsg)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -68,8 +68,13 @@ type listUsersResponse struct {
|
||||
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request query listUsersRequest true "查询参数"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=user.listUsersResponse} "用户列表"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/users [get]
|
||||
func ListUsers(c *gin.Context) {
|
||||
var req listUsersRequest
|
||||
@@ -129,9 +134,15 @@ type updateUserStatusRequest struct {
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param id path int true "用户ID"
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Param request body updateUserStatusRequest true "状态参数"
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "更新成功"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限或尝试禁用管理员"
|
||||
// @Failure 404 {object} util.ResponseAny "用户不存在"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/admin/users/{id}/status [put]
|
||||
func UpdateUserStatus(c *gin.Context) {
|
||||
var req updateUserStatusRequest
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
/*
|
||||
Copyright 2026 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 user
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
"github.com/shopspring/decimal"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/users", ListUsers)
|
||||
adminGroup.PUT("/users/:id/status", UpdateUserStatus)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestListUsers(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed users
|
||||
users := []model.User{
|
||||
{
|
||||
ID: 1001,
|
||||
Username: "alice",
|
||||
Nickname: "Alice Nickname",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
AvailableBalance: decimal.NewFromFloat(100.0),
|
||||
LastLoginAt: time.Now(),
|
||||
SignKey: "alice_sign_key",
|
||||
},
|
||||
{
|
||||
ID: 1002,
|
||||
Username: "bob",
|
||||
Nickname: "Bob Nickname",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
AvailableBalance: decimal.NewFromFloat(50.0),
|
||||
LastLoginAt: time.Now(),
|
||||
SignKey: "bob_sign_key",
|
||||
},
|
||||
{
|
||||
ID: 1003,
|
||||
Username: "charlie",
|
||||
Nickname: "Charlie Nickname",
|
||||
IsActive: false,
|
||||
IsAdmin: true,
|
||||
AvailableBalance: decimal.NewFromFloat(9999.0),
|
||||
LastLoginAt: time.Now(),
|
||||
SignKey: "charlie_sign_key",
|
||||
},
|
||||
}
|
||||
|
||||
for _, u := range users {
|
||||
if err := dbConn.Create(&u).Error; err != nil {
|
||||
t.Fatalf("failed to seed user: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true, SignKey: "charlie_sign_key"}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("basic pagination list", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=2", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("failed to parse response: %v", err)
|
||||
}
|
||||
|
||||
// Parse data map to our structure
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
if err := json.Unmarshal(dataBytes, &listResp); err != nil {
|
||||
t.Fatalf("failed to parse list response: %v", err)
|
||||
}
|
||||
|
||||
if len(listResp.Users) != 2 {
|
||||
t.Errorf("expected 2 users, got %d", len(listResp.Users))
|
||||
}
|
||||
if listResp.Total != 3 {
|
||||
t.Errorf("expected total 3, got %d", listResp.Total)
|
||||
}
|
||||
// Ordered by ID DESC
|
||||
if listResp.Users[0].ID != 1003 || listResp.Users[1].ID != 1002 {
|
||||
t.Errorf("expected ordered DESC, got first ID %d, second ID %d", listResp.Users[0].ID, listResp.Users[1].ID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by user_id", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&user_id=1001", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
json.Unmarshal(dataBytes, &listResp)
|
||||
|
||||
if len(listResp.Users) != 1 || listResp.Users[0].ID != 1001 {
|
||||
t.Errorf("expected 1 user with ID 1001, got total %d", len(listResp.Users))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by username prefix", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&username=bo", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
json.Unmarshal(dataBytes, &listResp)
|
||||
|
||||
if len(listResp.Users) != 1 || listResp.Users[0].Username != "bob" {
|
||||
t.Errorf("expected bob, got %v", listResp.Users)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid pagination parameter", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=0&page_size=10", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateUserStatus(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed users
|
||||
regularUser := model.User{
|
||||
ID: 1001,
|
||||
Username: "alice",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
SignKey: "alice_sign_key",
|
||||
}
|
||||
adminUser := model.User{
|
||||
ID: 1002,
|
||||
Username: "bob",
|
||||
IsActive: true,
|
||||
IsAdmin: true,
|
||||
SignKey: "bob_sign_key",
|
||||
}
|
||||
|
||||
dbConn.Create(®ularUser)
|
||||
dbConn.Create(&adminUser)
|
||||
|
||||
router := setupTestRouter(&adminUser)
|
||||
|
||||
t.Run("disable regular user successfully", func(t *testing.T) {
|
||||
payload := updateUserStatusRequest{IsActive: false}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1001/status", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// Verify DB status
|
||||
var u model.User
|
||||
dbConn.First(&u, 1001)
|
||||
if u.IsActive {
|
||||
t.Error("user should be deactivated in the database")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cannot disable admin user", func(t *testing.T) {
|
||||
payload := updateUserStatusRequest{IsActive: false}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1002/status", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Errorf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp util.ResponseAny
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
if resp.ErrorMsg != cannotDisable {
|
||||
t.Errorf("expected error message '%s', got '%s'", cannotDisable, resp.ErrorMsg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cannot enable/disable non-existent user", func(t *testing.T) {
|
||||
payload := updateUserStatusRequest{IsActive: false}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/9999/status", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -28,7 +28,10 @@ import (
|
||||
type PublicConfigResponse struct {
|
||||
UploadAllowedExtensions string `json:"upload_allowed_extensions"` // 允许上传的图片扩展名
|
||||
SiteName string `json:"site_name"` // 站点名称
|
||||
PasswordLoginEnabled bool `json:"password_login_enabled"` // 是否允许密码登录
|
||||
RegistrationEnabled bool `json:"registration_enabled"` // 是否允许注册
|
||||
PasswordRegisterEnabled bool `json:"password_register_enabled"` // 是否允许密码注册
|
||||
OIDCLoginEnabled bool `json:"oidc_login_enabled"` // 是否允许 OIDC 登录
|
||||
MaxAPIKeysPerUser int `json:"max_api_keys_per_user"` // 每个用户最大 API Key 数量
|
||||
}
|
||||
|
||||
@@ -62,6 +65,24 @@ func GetPublicConfig(c *gin.Context) {
|
||||
registrationEnabled = val
|
||||
}
|
||||
|
||||
// 3.1 password_login_enabled
|
||||
var passwordLoginEnabled bool
|
||||
if val, err := model.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled); err == nil {
|
||||
passwordLoginEnabled = val
|
||||
}
|
||||
|
||||
// 3.2 password_register_enabled
|
||||
var passwordRegisterEnabled bool
|
||||
if val, err := model.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled); err == nil {
|
||||
passwordRegisterEnabled = val
|
||||
}
|
||||
|
||||
// 3.3 oidc_login_enabled
|
||||
var oidcLoginEnabled bool
|
||||
if val, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled); err == nil {
|
||||
oidcLoginEnabled = val
|
||||
}
|
||||
|
||||
// 4. max_api_keys_per_user
|
||||
var maxAPIKeys int
|
||||
if val, err := model.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
|
||||
@@ -71,7 +92,10 @@ func GetPublicConfig(c *gin.Context) {
|
||||
response := PublicConfigResponse{
|
||||
UploadAllowedExtensions: uploadExtensions,
|
||||
SiteName: siteName,
|
||||
PasswordLoginEnabled: passwordLoginEnabled,
|
||||
RegistrationEnabled: registrationEnabled,
|
||||
PasswordRegisterEnabled: passwordRegisterEnabled,
|
||||
OIDCLoginEnabled: oidcLoginEnabled,
|
||||
MaxAPIKeysPerUser: maxAPIKeys,
|
||||
}
|
||||
|
||||
|
||||
@@ -23,12 +23,12 @@ import (
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
// Health godoc
|
||||
// Health 健康检查
|
||||
// @Summary 健康检查
|
||||
// @Description 检查服务是否正常运行
|
||||
// @Description 检查服务是否正常运行,可用于负载均衡存活探测
|
||||
// @Tags health
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "服务正常"
|
||||
// @Router /api/v1/health [get]
|
||||
func Health(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
|
||||
@@ -19,6 +19,7 @@ package oauth
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"strings"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
@@ -35,14 +36,19 @@ func init() {
|
||||
|
||||
if cfg.Issuer != "" {
|
||||
ctx := context.Background()
|
||||
provider, err := oidc.NewProvider(ctx, cfg.Issuer)
|
||||
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
|
||||
issuer := strings.TrimSuffix(strings.TrimSpace(cfg.Issuer), "/")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
|
||||
|
||||
provider, err := oidc.NewProvider(ctx, 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)
|
||||
log.Printf("[OAuth] OIDC Provider 初始化成功: %s", issuer)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ limitations under the License.
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -30,3 +31,29 @@ const (
|
||||
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
||||
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
||||
)
|
||||
|
||||
const (
|
||||
OAuthPurposeLogin = "login"
|
||||
OAuthPurposeBind = "bind"
|
||||
)
|
||||
|
||||
type oauthStatePayload struct {
|
||||
SourceName string `json:"source_name"`
|
||||
Purpose string `json:"purpose"`
|
||||
}
|
||||
|
||||
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
|
||||
var payload oauthStatePayload
|
||||
if err := json.Unmarshal([]byte(value), &payload); err != nil {
|
||||
return oauthStatePayload{}, err
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
@@ -55,60 +55,49 @@ func doOAuth(ctx context.Context, code string, nonce string) (*model.User, error
|
||||
|
||||
var userInfo model.OAuthUserInfo
|
||||
|
||||
if config.Config.App.Env == "development" && code == "dev_mock_code" {
|
||||
userInfo = model.OAuthUserInfo{
|
||||
Id: 999999,
|
||||
Username: "dev_user",
|
||||
Name: "Developer User",
|
||||
Active: true,
|
||||
AvatarUrl: "https://linux.do/user_avatar/linux.do/system/45/1_2.png",
|
||||
TrustLevel: 3,
|
||||
}
|
||||
} else {
|
||||
// 使用授权码换取 Token
|
||||
token, err := oauthConf.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
// 使用授权码换取 Token
|
||||
token, err := oauthConf.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
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 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()
|
||||
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
|
||||
}
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -119,7 +108,7 @@ func doOAuth(ctx context.Context, code string, nonce string) (*model.User, error
|
||||
}
|
||||
|
||||
var user model.User
|
||||
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
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 {
|
||||
// 存在冲突 -> 将占用者改名并注销
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+43
-115
@@ -17,104 +17,15 @@ 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/config"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
"github.com/shopspring/decimal"
|
||||
)
|
||||
|
||||
// GetLoginURL godoc
|
||||
// @Summary 获取登录地址
|
||||
// @Description 生成 OAuth 登录 URL,前端跳转至该地址完成授权
|
||||
// @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 config.Config.App.Env == "development" {
|
||||
authURL = fmt.Sprintf("%s/login?code=dev_mock_code&state=%s", config.Config.App.FrontendURL, state)
|
||||
} else 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
|
||||
// @Summary OAuth 回调
|
||||
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立用户会话
|
||||
// @Tags oauth
|
||||
// @Accept json
|
||||
// @Param request body CallbackRequest true "回调请求参数"
|
||||
// @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"`
|
||||
@@ -135,46 +46,63 @@ type BasicUserInfo struct {
|
||||
DailyLimit *int64 `json:"daily_limit"`
|
||||
}
|
||||
|
||||
// UserInfo godoc
|
||||
func BuildBasicUserInfo(user *model.User) BasicUserInfo {
|
||||
return 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,
|
||||
IsAdmin: user.IsAdmin,
|
||||
RemainQuota: decimal.NewFromInt(-1),
|
||||
PayLevel: "Free",
|
||||
DailyLimit: nil,
|
||||
}
|
||||
}
|
||||
|
||||
// UserInfo 获取当前登录用户信息
|
||||
// @Summary 获取当前登录用户信息
|
||||
// @Description 返回当前登录用户的基本信息及余额数据,需要登录
|
||||
// @Description 返回当前登录用户的基本信息及余额数据,需要登录。包括用户 ID、用户名、信任等级、各类余额信息等。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.BasicUserInfo} "用户信息"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Router /api/v1/oauth/user-info [get]
|
||||
func UserInfo(c *gin.Context) {
|
||||
user, _ := util.GetFromContext[*model.User](c, UserObjKey)
|
||||
|
||||
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,
|
||||
IsAdmin: user.IsAdmin,
|
||||
RemainQuota: decimal.NewFromInt(-1),
|
||||
PayLevel: "Free",
|
||||
DailyLimit: nil,
|
||||
}),
|
||||
util.OK(BuildBasicUserInfo(user)),
|
||||
)
|
||||
}
|
||||
|
||||
// Logout godoc
|
||||
// @Summary 退出登录
|
||||
// @Description 清除当前用户的登录会话,完成退出
|
||||
// GetLoginURL 获取登录地址
|
||||
// @Summary 获取登录地址
|
||||
// @Description 生成 OAuth 登录 URL,前端跳转至该地址完成授权。返回的 URL 中包含 state 参数用于 CSRF 防护。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "OAuth 登录 URL"
|
||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或内部错误"
|
||||
// @Router /api/v1/oauth/login [get]
|
||||
|
||||
// Logout 退出登录
|
||||
// @Summary 退出登录
|
||||
// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "退出成功"
|
||||
// @Failure 500 {object} util.ResponseAny "Session 清除失败"
|
||||
// @Router /api/v1/oauth/logout [get]
|
||||
func Logout(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
|
||||
@@ -0,0 +1,611 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"strconv"
|
||||
|
||||
"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/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/util"
|
||||
"golang.org/x/oauth2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// AuthSourceView 登录源展示信息
|
||||
type AuthSourceView struct {
|
||||
ID uint64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IconURL string `json:"icon_url"`
|
||||
ClientSecretConfigured bool `json:"client_secret_configured"`
|
||||
}
|
||||
|
||||
// OAuthAuthorizeResponse 授权 URL 响应
|
||||
type OAuthAuthorizeResponse struct {
|
||||
AuthorizeURL string `json:"authorize_url"`
|
||||
}
|
||||
|
||||
// OAuthCallbackResult 回调处理结果
|
||||
type OAuthCallbackResult struct {
|
||||
Status string `json:"status"`
|
||||
User *BasicUserInfo `json:"user,omitempty"`
|
||||
}
|
||||
|
||||
// CallbackRequest OAuth 回调请求参数
|
||||
type CallbackRequest struct {
|
||||
State string `json:"state" binding:"required"`
|
||||
Code string `json:"code" binding:"required"`
|
||||
}
|
||||
|
||||
func defaultAuthSource() *model.AuthSource {
|
||||
if config.Config.OAuth2.ClientID == "" || config.Config.OAuth2.RedirectURI == "" {
|
||||
return nil
|
||||
}
|
||||
source := &model.AuthSource{
|
||||
Name: "default",
|
||||
Type: model.AuthSourceTypeOIDC,
|
||||
DisplayName: "默认认证源",
|
||||
IsActive: true,
|
||||
ClientID: config.Config.OAuth2.ClientID,
|
||||
ClientSecret: config.Config.OAuth2.ClientSecret,
|
||||
OpenIDDiscoveryURL: config.Config.OAuth2.Issuer,
|
||||
}
|
||||
if source.DisplayName == "" {
|
||||
source.DisplayName = "默认认证源"
|
||||
}
|
||||
return source
|
||||
}
|
||||
|
||||
func resolveAuthSource(sourceName string) (*model.AuthSource, error) {
|
||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||
if name == "" || name == "default" {
|
||||
source := defaultAuthSource()
|
||||
if source == nil {
|
||||
return nil, errors.New("默认认证源未配置")
|
||||
}
|
||||
return source, nil
|
||||
}
|
||||
return model.GetAuthSourceByName(name)
|
||||
}
|
||||
|
||||
func activeLoginSources() []AuthSourceView {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyOIDCLoginEnabled)
|
||||
if err == nil && !enabled {
|
||||
return nil
|
||||
}
|
||||
|
||||
sources := make([]AuthSourceView, 0, 4)
|
||||
if source := defaultAuthSource(); source != nil {
|
||||
source.Sanitize()
|
||||
sources = append(sources, AuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
IsActive: source.IsActive,
|
||||
IconURL: source.IconURL,
|
||||
ClientSecretConfigured: source.ClientSecretConfigured,
|
||||
})
|
||||
}
|
||||
|
||||
dbSources, err := model.GetActiveAuthSources()
|
||||
if err != nil {
|
||||
return sources
|
||||
}
|
||||
for _, source := range dbSources {
|
||||
sources = append(sources, AuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
IsActive: source.IsActive,
|
||||
IconURL: source.IconURL,
|
||||
ClientSecretConfigured: source.ClientSecretConfigured,
|
||||
})
|
||||
}
|
||||
return sources
|
||||
}
|
||||
|
||||
func frontendLoginRedirectURL() string {
|
||||
if config.Config.App.FrontendURL != "" {
|
||||
return strings.TrimRight(config.Config.App.FrontendURL, "/") + "/login"
|
||||
}
|
||||
if config.Config.OAuth2.RedirectURI != "" {
|
||||
return config.Config.OAuth2.RedirectURI
|
||||
}
|
||||
return "/login"
|
||||
}
|
||||
|
||||
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
|
||||
if source == nil {
|
||||
return nil, nil, errors.New("认证源不能为空")
|
||||
}
|
||||
|
||||
if source.Name == "default" {
|
||||
scopes := []string{"profile", "email"}
|
||||
if oidcVerifier != nil {
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
return &oauth2.Config{
|
||||
ClientID: config.Config.OAuth2.ClientID,
|
||||
ClientSecret: config.Config.OAuth2.ClientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: scopes,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: config.Config.OAuth2.AuthorizationEndpoint,
|
||||
TokenURL: config.Config.OAuth2.TokenEndpoint,
|
||||
AuthStyle: oauth2.AuthStyleAutoDetect,
|
||||
},
|
||||
}, oidcVerifier, nil
|
||||
}
|
||||
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return nil, nil, errors.New("OIDC 认证源必须配置 Discovery URL")
|
||||
}
|
||||
|
||||
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
|
||||
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
|
||||
|
||||
provider, err := oidc.NewProvider(ctx, issuer)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
|
||||
scopes := strings.Fields(source.Scopes)
|
||||
if len(scopes) == 0 {
|
||||
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
|
||||
}
|
||||
if !containsScope(scopes, oidc.ScopeOpenID) {
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
|
||||
return &oauth2.Config{
|
||||
ClientID: source.ClientID,
|
||||
ClientSecret: source.ClientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: scopes,
|
||||
Endpoint: provider.Endpoint(),
|
||||
}, verifier, nil
|
||||
}
|
||||
|
||||
func containsScope(scopes []string, scope string) bool {
|
||||
for _, item := range scopes {
|
||||
if item == scope {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func setLoginSession(c *gin.Context, user *model.User) error {
|
||||
session := sessions.Default(c)
|
||||
session.Set(UserIDKey, user.ID)
|
||||
session.Set(UserNameKey, user.Username)
|
||||
return session.Save()
|
||||
}
|
||||
|
||||
func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
candidate := strings.TrimSpace(base)
|
||||
if candidate == "" {
|
||||
candidate = "user"
|
||||
}
|
||||
for i := 0; i < 1000; i++ {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", candidate).Count(&count).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
if count == 0 {
|
||||
return candidate, nil
|
||||
}
|
||||
candidate = fmt.Sprintf("%s-%d", base, i+1)
|
||||
}
|
||||
return "", errors.New("无法生成可用用户名")
|
||||
}
|
||||
|
||||
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
|
||||
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token, err := authConfig.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
userInfo := &model.OAuthUserInfo{Active: true}
|
||||
if verifier != nil {
|
||||
if rawIDToken, ok := token.Extra("id_token").(string); ok {
|
||||
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
|
||||
if verifyErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
|
||||
}
|
||||
if nonce != "" && idToken.Nonce != nonce {
|
||||
return nil, errors.New(NonceMismatch)
|
||||
}
|
||||
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
|
||||
return nil, claimsErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
|
||||
userInfo.Username = userInfo.PreferredUsername
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Email != "" {
|
||||
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Sub != "" {
|
||||
userInfo.Username = userInfo.Sub
|
||||
}
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
|
||||
if userInfo.Username == "" {
|
||||
client := authConfig.Client(ctx, token)
|
||||
userEndpoint := config.Config.OAuth2.UserEndpoint
|
||||
if source.Name != "default" {
|
||||
userEndpoint = ""
|
||||
}
|
||||
if userEndpoint != "" {
|
||||
resp, httpErr := client.Get(userEndpoint)
|
||||
if httpErr != nil {
|
||||
return nil, httpErr
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
responseData, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
return nil, readErr
|
||||
}
|
||||
if unmarshalErr := json.Unmarshal(responseData, userInfo); unmarshalErr != nil {
|
||||
return nil, unmarshalErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return userInfo, nil
|
||||
}
|
||||
|
||||
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
|
||||
userInfo.Username = strings.TrimSpace(userInfo.Username)
|
||||
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
|
||||
userInfo.Email = strings.TrimSpace(userInfo.Email)
|
||||
userInfo.Name = strings.TrimSpace(userInfo.Name)
|
||||
userInfo.AvatarUrl = strings.TrimSpace(userInfo.AvatarUrl)
|
||||
|
||||
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
|
||||
userInfo.Username = userInfo.PreferredUsername
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Email != "" {
|
||||
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
|
||||
}
|
||||
if userInfo.Username == "" && userInfo.Sub != "" {
|
||||
userInfo.Username = userInfo.Sub
|
||||
}
|
||||
if userInfo.Username == "" {
|
||||
return errors.New("无法从认证源获取用户名")
|
||||
}
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
if !userInfo.Active {
|
||||
userInfo.Active = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
|
||||
result := OAuthCallbackResult{Status: status}
|
||||
if user != nil {
|
||||
info := BuildBasicUserInfo(user)
|
||||
result.User = &info
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// GetLoginSources 获取可用登录源列表
|
||||
// @Summary 获取可用登录源
|
||||
// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Success 200 {object} util.ResponseAny{data=[]oauth.AuthSourceView} "登录源列表"
|
||||
// @Router /api/v1/oauth/sources [get]
|
||||
func GetLoginSources(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, util.OK(activeLoginSources()))
|
||||
}
|
||||
|
||||
// GetLoginURL 获取登录授权地址
|
||||
// @Summary 获取登录授权地址
|
||||
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Param source query string false "认证源名称,为空使用默认源"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||
// @Failure 400 {object} util.ResponseAny "认证源不存在或未配置"
|
||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
||||
// @Router /api/v1/oauth/login [get]
|
||||
func GetLoginURL(c *gin.Context) {
|
||||
source, err := resolveAuthSource(c.Query("source"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: source.Name,
|
||||
Purpose: OAuthPurposeLogin,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
|
||||
}
|
||||
|
||||
func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
|
||||
authConfig, verifier, err := buildOAuthConfig(ctx, source, frontendLoginRedirectURL())
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if verifier != nil {
|
||||
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
|
||||
}
|
||||
return authConfig.AuthCodeURL(state), nil
|
||||
}
|
||||
|
||||
// Authorize 发起指定认证源授权
|
||||
// @Summary 发起指定认证源授权
|
||||
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Param source path string true "认证源名称"
|
||||
// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL"
|
||||
// @Failure 400 {object} util.ResponseAny "认证源不存在或未启用"
|
||||
// @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败"
|
||||
// @Router /api/v1/oauth/{source}/authorize [get]
|
||||
func Authorize(c *gin.Context) {
|
||||
source, err := resolveAuthSource(c.Param("source"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if !source.IsActive {
|
||||
c.JSON(http.StatusBadRequest, util.Err("认证源未启用"))
|
||||
return
|
||||
}
|
||||
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
|
||||
if purpose != OAuthPurposeBind {
|
||||
purpose = OAuthPurposeLogin
|
||||
}
|
||||
state := uuid.NewString()
|
||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||
SourceName: source.Name,
|
||||
Purpose: purpose,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
|
||||
}
|
||||
|
||||
// Callback OAuth 回调处理
|
||||
// @Summary OAuth 回调处理
|
||||
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
|
||||
// @Tags oauth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body oauth.CallbackRequest true "回调请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.OAuthCallbackResult} "登录或绑定成功"
|
||||
// @Failure 400 {object} util.ResponseAny "state 无效、参数错误或认证源错误"
|
||||
// @Failure 401 {object} util.ResponseAny "绑定场景未登录"
|
||||
// @Failure 500 {object} util.ResponseAny "OAuth 认证失败或内部错误"
|
||||
// @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()
|
||||
stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(InvalidState))
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Del(ctx, stateKey)
|
||||
|
||||
payload, err := decodeOAuthStatePayload(payloadRaw)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
source, err := resolveAuthSource(payload.SourceName)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, frontendLoginRedirectURL())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := normalizeOAuthUserInfo(userInfo); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if userInfo.Sub == "" {
|
||||
userInfo.Sub = userInfo.Username
|
||||
}
|
||||
|
||||
var user model.User
|
||||
if payload.Purpose == OAuthPurposeBind {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
c.JSON(http.StatusUnauthorized, util.Err(common.UnAuthorized))
|
||||
return
|
||||
}
|
||||
if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := model.BindExternalAccount(&model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
|
||||
c.JSON(http.StatusOK, util.OK(buildCallbackResult(&user, "bound")))
|
||||
return
|
||||
}
|
||||
|
||||
account, err := model.FindExternalAccount(source.ID, userInfo.Sub)
|
||||
switch {
|
||||
case err == nil:
|
||||
if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
|
||||
if uniqueErr != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error()))
|
||||
return
|
||||
}
|
||||
user = model.User{
|
||||
Username: username,
|
||||
Nickname: userInfo.Name,
|
||||
AvatarUrl: userInfo.AvatarUrl,
|
||||
TrustLevel: userInfo.TrustLevel,
|
||||
SignKey: util.GenerateUniqueIDSimple(),
|
||||
IsActive: true,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
if err := db.DB(ctx).Create(&user).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := model.BindExternalAccount(&model.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
default:
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
|
||||
if err := setLoginSession(c, &user); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(buildCallbackResult(&user, "logged_in")))
|
||||
}
|
||||
|
||||
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
|
||||
// @Summary 获取外部帐号列表
|
||||
// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=[]model.ExternalAccountView} "外部帐号列表"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 500 {object} util.ResponseAny "内部错误"
|
||||
// @Router /api/v1/oauth/external-accounts [get]
|
||||
func ListExternalAccounts(c *gin.Context) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
accounts, err := model.ListExternalAccountsByUserID(userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OK(accounts))
|
||||
}
|
||||
|
||||
// DeleteExternalAccount 解除外部帐号绑定
|
||||
// @Summary 解除外部帐号绑定
|
||||
// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
|
||||
// @Tags oauth
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "外部帐号绑定记录 ID"
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "解除绑定成功"
|
||||
// @Failure 400 {object} util.ResponseAny "ID 无效或解除失败"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
|
||||
func DeleteExternalAccount(c *gin.Context) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
c.JSON(http.StatusUnauthorized, util.Err(common.UnAuthorized))
|
||||
return
|
||||
}
|
||||
rawID := strings.TrimSpace(c.Param("id"))
|
||||
id, err := strconv.ParseUint(rawID, 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
c.JSON(http.StatusBadRequest, util.Err("绑定记录 ID 无效"))
|
||||
return
|
||||
}
|
||||
if err := model.DeleteExternalAccountForUser(id, userID); err != nil {
|
||||
c.JSON(http.StatusBadRequest, util.Err(err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, util.OKNil())
|
||||
}
|
||||
@@ -28,11 +28,16 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// ServeFileByID serves an uploaded file by its ID
|
||||
// ServeFileByID 根据 ID 获取并提供已上传的文件
|
||||
// @Summary 获取已上传文件
|
||||
// @Description 根据文件 ID 获取并提供已上传的临时或正式文件,若配置了缓存则优先走本地缓存,否则从 S3 等后端存储读取并流式返回
|
||||
// @Tags upload
|
||||
// @Produce octet-stream
|
||||
// @Param id path string true "Upload ID"
|
||||
// @Success 200
|
||||
// @Param id path string true "文件 ID"
|
||||
// @Success 200 {file} file "成功获取文件内容"
|
||||
// @Failure 400 {object} util.ResponseAny "文件 ID 格式错误"
|
||||
// @Failure 404 {object} util.ResponseAny "文件未找到"
|
||||
// @Failure 500 {object} util.ResponseAny "服务内部错误"
|
||||
// @Router /f/{id} [get]
|
||||
func ServeFileByID(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/apps/oauth"
|
||||
"github.com/linux-do/credit/internal/common"
|
||||
"github.com/linux-do/credit/internal/common/bind"
|
||||
"github.com/linux-do/credit/internal/common/response"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type registerRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Nickname string `json:"nickname"`
|
||||
DisplayName string `json:"display_name"`
|
||||
}
|
||||
|
||||
func isPasswordLoginEnabled() bool {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isPasswordRegisterEnabled() bool {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordRegisterEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func isRegistrationEnabled() bool {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyRegistrationEnabled)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
func setLoginSession(c *gin.Context, user *model.User) error {
|
||||
session := sessions.Default(c)
|
||||
session.Set(oauth.UserIDKey, user.ID)
|
||||
session.Set(oauth.UserNameKey, user.Username)
|
||||
if err := session.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Login 用户密码登录
|
||||
// @Summary 用户密码登录
|
||||
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.loginRequest true "登录请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.BasicUserInfo} "登录成功,返回用户信息"
|
||||
// @Failure 400 {object} util.ResponseAny "用户名或密码错误、帐号已禁用等"
|
||||
// @Failure 500 {object} util.ResponseAny "服务内部错误"
|
||||
// @Router /api/v1/user/login [post]
|
||||
func Login(c *gin.Context) {
|
||||
if !isPasswordLoginEnabled() {
|
||||
response.RespondFailure(c, "管理员关闭了密码登录")
|
||||
return
|
||||
}
|
||||
var req loginRequest
|
||||
if !bind.JSON(c, &req) {
|
||||
return
|
||||
}
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
if req.Username == "" || req.Password == "" {
|
||||
response.RespondFailure(c, "无效的参数")
|
||||
return
|
||||
}
|
||||
|
||||
var user model.User
|
||||
ctx := c.Request.Context()
|
||||
if err := db.DB(ctx).Where("username = ?", req.Username).First(&user).Error; err != nil {
|
||||
response.RespondFailure(c, "用户名或密码错误")
|
||||
return
|
||||
}
|
||||
if !user.IsActive {
|
||||
response.RespondFailure(c, common.BannedAccount)
|
||||
return
|
||||
}
|
||||
if !user.CheckPassword(req.Password) {
|
||||
response.RespondFailure(c, "用户名或密码错误")
|
||||
return
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := setLoginSession(c, &user); err != nil {
|
||||
response.RespondFailure(c, "无法保存会话信息,请重试")
|
||||
return
|
||||
}
|
||||
|
||||
response.RespondSuccess(c, oauth.BuildBasicUserInfo(&user))
|
||||
}
|
||||
|
||||
// Register 用户注册
|
||||
// @Summary 用户注册
|
||||
// @Description 使用用户名和密码注册新账号,注册成功后自动登录并建立 Session。密码长度不能少于 8 位。
|
||||
// @Tags user
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body user.registerRequest true "注册请求参数"
|
||||
// @Success 200 {object} util.ResponseAny{data=oauth.BasicUserInfo} "注册并登录成功,返回用户信息"
|
||||
// @Failure 400 {object} util.ResponseAny "参数错误、用户名已存在或注册已关闭"
|
||||
// @Failure 500 {object} util.ResponseAny "服务内部错误"
|
||||
// @Router /api/v1/user/register [post]
|
||||
func Register(c *gin.Context) {
|
||||
if !isRegistrationEnabled() || !isPasswordRegisterEnabled() {
|
||||
response.RespondFailure(c, "管理员关闭了注册")
|
||||
return
|
||||
}
|
||||
|
||||
var req registerRequest
|
||||
if !bind.JSON(c, &req) {
|
||||
return
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
req.Password = strings.TrimSpace(req.Password)
|
||||
req.Nickname = strings.TrimSpace(req.Nickname)
|
||||
req.DisplayName = strings.TrimSpace(req.DisplayName)
|
||||
|
||||
if req.Username == "" || req.Password == "" {
|
||||
response.RespondFailure(c, "无效的参数")
|
||||
return
|
||||
}
|
||||
if len(req.Password) < 8 {
|
||||
response.RespondFailure(c, "密码长度不能少于 8 位")
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
if count > 0 {
|
||||
response.RespondFailure(c, "用户名已存在")
|
||||
return
|
||||
}
|
||||
|
||||
user := model.User{
|
||||
Username: req.Username,
|
||||
Nickname: req.Nickname,
|
||||
AvatarUrl: "",
|
||||
TrustLevel: model.TrustLevelNewUser,
|
||||
PayScore: 0,
|
||||
SignKey: util.GenerateUniqueIDSimple(),
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = req.DisplayName
|
||||
}
|
||||
if user.Nickname == "" {
|
||||
user.Nickname = req.Username
|
||||
}
|
||||
if err := user.SetPassword(req.Password); err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Create(&user).Error; err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := setLoginSession(c, &user); err != nil {
|
||||
response.RespondFailure(c, "无法保存会话信息,请重试")
|
||||
return
|
||||
}
|
||||
|
||||
response.RespondSuccess(c, oauth.BuildBasicUserInfo(&user))
|
||||
}
|
||||
|
||||
// Logout 用户退出登录
|
||||
// @Summary 用户退出登录
|
||||
// @Description 清除用户登录 Session,完成退出
|
||||
// @Tags user
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=string} "退出成功"
|
||||
// @Failure 500 {object} util.ResponseAny "Session 清除失败"
|
||||
// @Router /api/v1/user/logout [get]
|
||||
func Logout(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Options(util.GetSessionOptions(-1))
|
||||
session.Clear()
|
||||
if err := session.Save(); err != nil {
|
||||
response.RespondFailure(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.RespondSuccessMessage(c, "")
|
||||
}
|
||||
Reference in New Issue
Block a user