[优化] response 结构调整

This commit is contained in:
ryan
2026-06-05 11:42:25 +08:00
parent 4dc4c745c8
commit 43e8ef154f
29 changed files with 578 additions and 605 deletions
+1 -1
View File
@@ -26,7 +26,7 @@ sidebar: false
### 变更 ### 变更
- 提取并新增 `common/response` 通用响应子包,替换 `controller/response.go` 中的响应细节,重构所有中间件中硬编码的 `c.JSON` 响应,统一 API 响应格式并隔离 Gin 依赖。 - 提取并新增 `common/response` 通用响应子包与 `controller/bind` 参数解析绑定子包,重构所有控制器与中间件中硬编码的 `c.JSON` 及参数绑定逻辑,彻底隔离 Gin 依赖并避免变量遮蔽冲突。
- 将 `usage.md` 改名为 `proxy-config.md`(新建反代配置),重新梳理大纲结构,专注于如何从导入/申请证书开始,一步步新增并发布代理路由规则,并同步更新全部导航与文档引用链接 - 将 `usage.md` 改名为 `proxy-config.md`(新建反代配置),重新梳理大纲结构,专注于如何从导入/申请证书开始,一步步新增并发布代理路由规则,并同步更新全部导航与文档引用链接
- WAF 白名单调整为准入名单语义:存在白名单规则时,未命中白名单的请求会被拦截 - WAF 白名单调整为准入名单语义:存在白名单规则时,未命中白名单的请求会被拦截
@@ -1,53 +0,0 @@
# 提取公共 Response 包以支持 Middleware 统一响应实现计划
---
## 1. 目标与背景 (Goal & Context)
* **需求背景**:
原本 `response.go` 放在 `controller` 目录下。由于 `middleware` 处于 `controller` 上游,并且和 `controller` 跨包,`middleware` 无法直接调用 `controller` 的响应逻辑。这导致 `middleware` 中存在许多手写的、格式硬编码的 `c.JSON` 调用。
如果直接将 `response.go` 放入已有的 `common` 根包,会引入 `github.com/gin-gonic/gin` 依赖,导致依赖 `common` 的底层 `service` 和 `model` 也受到 `gin` 框架的依赖污染。
* **开发范围 (Scope)**:
* 在 `openflare_server/common/response` 下新建 `response.go` 子包(`package response`),存放通用的 HTTP 响应逻辑。
* 将 `controller/response.go` 中的响应逻辑移植到 `common/response/response.go`。
* 在 `controller/response.go` 中保留参数解析逻辑,并作为代理调用 `common/response` 中的方法,从而实现对 controller 内 140+ 处现有调用的**零改动**。
* 修改 `middleware` 包下的所有 `c.JSON` 手写响应,统一采用 `common/response` 的方法。
## 2. 设计与决策 (Design & Decisions)
* **核心对象/数据模型**:无需改动任何数据库或数据模型。
* **API 与鉴权设计**:不改变任何公开 API 接口路由与现有的成功/失败响应 JSON 格式。
* **设计决策权衡**:
* **方案 A:直接把 response.go 放入 common 根包**
* 缺点:导致原本不应该感知 Web 传输协议的 `service`、`model`、`job` 包间接依赖了 `github.com/gin-gonic/gin` 框架。
* **方案 B:在 common 下建子包 `common/response`**
* 优点:满足 `common` 归类的直觉,又保持了包的物理依赖隔离。业务底层不受 `gin` 污染,而 `controller` 和 `middleware` 这类传输层可引入该包进行代码复用。**(采用此方案)**
## 3. 具体修改文件清单 (Proposed Changes)
### 后端 Server
* #### [NEW] [response.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/common/response/response.go)
* 职责:通用的 Gin 响应工具包,包括 `RespondSuccess`、`RespondSuccessWithExtras`、`RespondSuccessMessage`、`RespondFailure`、`RespondBadRequest`、`RespondUnauthorized` 和 `RespondForbidden`。
* #### [MODIFY] [response.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/controller/response.go)
* 职责:删去具体的 HTTP 响应渲染实现,以代理方式调用 `common/response` 包中导出的函数,保持 controller 包内现有调用的向后兼容;保留原有的 `decodeJSONBody`、`decodeOptionalJSONBody`、`parseIDParam`、`parseIDParamByName`、`bindJSON` 参数绑定解析逻辑。
* #### [MODIFY] [agent-auth.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/middleware/agent-auth.go)
* 职责:替换 `c.JSON(http.StatusUnauthorized, ...)` 为使用 `response.RespondUnauthorized(...)` 渲染统一响应。
* #### [MODIFY] [auth.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/middleware/auth.go)
* 职责:替换 `c.JSON` 相关的未授权和失败响应为使用 `response` 包方法。
* #### [MODIFY] [jwt.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/middleware/jwt.go)
* 职责:替换 JWT 未授权的回调响应为使用 `response` 统一格式。
* #### [MODIFY] [relay-auth.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/middleware/relay-auth.go)
* 职责:替换未授权和 StatusForbidden 响应为使用 `response` 方法。
* #### [MODIFY] [tunnel-auth.go](file:///Users/ryan/DEV/Go/OpenFlare/openflare_server/middleware/tunnel-auth.go)
* 职责:替换未授权和 StatusForbidden 响应为使用 `response` 方法。
---
## 4. 验证计划 (Verification Plan)
### 自动化单元测试
* 运行项目已有测试验证重构是否影响 API 连通性与响应结构:
```bash
go test -v ./router/...
```
```bash
go test -v ./service/...
```
+15 -13
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"strconv" "strconv"
@@ -25,10 +27,10 @@ import (
func GetAccessLogs(c *gin.Context) { func GetAccessLogs(c *gin.Context) {
logs, err := service.ListAccessLogs(readAccessLogQuery(c)) logs, err := service.ListAccessLogs(readAccessLogQuery(c))
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, logs) response.RespondSuccess(c, logs)
} }
// GetFoldedAccessLogs godoc // GetFoldedAccessLogs godoc
@@ -52,10 +54,10 @@ func GetFoldedAccessLogs(c *gin.Context) {
query.FoldMinutes = readQueryInt(c, "fold_minutes") query.FoldMinutes = readQueryInt(c, "fold_minutes")
logs, err := service.ListFoldedAccessLogs(query) logs, err := service.ListFoldedAccessLogs(query)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, logs) response.RespondSuccess(c, logs)
} }
// GetFoldedAccessLogIPs godoc // GetFoldedAccessLogIPs godoc
@@ -89,10 +91,10 @@ func GetFoldedAccessLogIPs(c *gin.Context) {
SortOrder: c.Query("sort_order"), SortOrder: c.Query("sort_order"),
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// GetAccessLogIPSummaries godoc // GetAccessLogIPSummaries godoc
@@ -120,10 +122,10 @@ func GetAccessLogIPSummaries(c *gin.Context) {
SortOrder: c.Query("sort_order"), SortOrder: c.Query("sort_order"),
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// GetAccessLogIPTrend godoc // GetAccessLogIPTrend godoc
@@ -147,10 +149,10 @@ func GetAccessLogIPTrend(c *gin.Context) {
BucketMinutes: readQueryInt(c, "bucket_minutes"), BucketMinutes: readQueryInt(c, "bucket_minutes"),
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// CleanupAccessLogs godoc // CleanupAccessLogs godoc
@@ -163,15 +165,15 @@ func GetAccessLogIPTrend(c *gin.Context) {
// @Router /api/access-logs/cleanup [post] // @Router /api/access-logs/cleanup [post]
func CleanupAccessLogs(c *gin.Context) { func CleanupAccessLogs(c *gin.Context) {
var input service.AccessLogCleanupInput var input service.AccessLogCleanupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
result, err := service.CleanupAccessLogs(input) result, err := service.CleanupAccessLogs(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
func readAccessLogQuery(c *gin.Context) service.AccessLogQuery { func readAccessLogQuery(c *gin.Context) service.AccessLogQuery {
+3 -2
View File
@@ -1,6 +1,7 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/model" "openflare/model"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -16,8 +17,8 @@ import (
func GetDefaultAcmeAccount(c *gin.Context) { func GetDefaultAcmeAccount(c *gin.Context) {
account, err := model.GetDefaultAcmeAccount() account, err := model.GetDefaultAcmeAccount()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, account) response.RespondSuccess(c, account)
} }
+34 -32
View File
@@ -5,6 +5,8 @@ import (
"log/slog" "log/slog"
"net" "net"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"strconv" "strconv"
@@ -26,7 +28,7 @@ import (
// @Router /api/agent/nodes/register [post] // @Router /api/agent/nodes/register [post]
func AgentRegister(c *gin.Context) { func AgentRegister(c *gin.Context) {
var payload service.AgentNodePayload var payload service.AgentNodePayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
@@ -41,10 +43,10 @@ func AgentRegister(c *gin.Context) {
result, err = service.RegisterNodeWithDiscovery(payload) result, err = service.RegisterNodeWithDiscovery(payload)
} }
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// AgentHeartbeat godoc // AgentHeartbeat godoc
@@ -59,23 +61,23 @@ func AgentRegister(c *gin.Context) {
// @Router /api/agent/nodes/heartbeat [post] // @Router /api/agent/nodes/heartbeat [post]
func AgentHeartbeat(c *gin.Context) { func AgentHeartbeat(c *gin.Context) {
var payload service.AgentNodePayload var payload service.AgentNodePayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := c.Get("agent_node") authNode, ok := c.Get("agent_node")
if !ok { if !ok {
respondUnauthorized(c, "鏃犳潈杩涜姝ゆ搷浣滐紝Agent Token 鏃犳晥") response.RespondUnauthorized(c, "鏃犳潈杩涜姝ゆ搷浣滐紝Agent Token 鏃犳晥")
return return
} }
node, err := service.HeartbeatNode(authNode.(*model.Node), payload) node, err := service.HeartbeatNode(authNode.(*model.Node), payload)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessWithExtras(c, node.Node, gin.H{ response.RespondSuccessWithExtras(c, node.Node, gin.H{
"agent_settings": node.AgentSettings, "agent_settings": node.AgentSettings,
"active_config": node.ActiveConfig, "active_config": node.ActiveConfig,
"waf_ip_groups": node.WAFIPGroups, "waf_ip_groups": node.WAFIPGroups,
@@ -94,15 +96,15 @@ func AgentHeartbeat(c *gin.Context) {
// @Router /api/agent/waf/ip-groups/sync [post] // @Router /api/agent/waf/ip-groups/sync [post]
func AgentSyncWAFIPGroups(c *gin.Context) { func AgentSyncWAFIPGroups(c *gin.Context) {
var input service.AgentWAFIPGroupSyncInput var input service.AgentWAFIPGroupSyncInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
result, err := service.SyncWAFIPGroupsForAgent(input) result, err := service.SyncWAFIPGroupsForAgent(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// AgentGetActiveConfig godoc // AgentGetActiveConfig godoc
@@ -115,7 +117,7 @@ func AgentSyncWAFIPGroups(c *gin.Context) {
func AgentGetActiveConfig(c *gin.Context) { func AgentGetActiveConfig(c *gin.Context) {
authNode, ok := c.Get("agent_node") authNode, ok := c.Get("agent_node")
if !ok { if !ok {
respondUnauthorized(c, "Node object missing from context") response.RespondUnauthorized(c, "Node object missing from context")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
@@ -123,19 +125,19 @@ func AgentGetActiveConfig(c *gin.Context) {
if node.NodeType == "tunnel_client" { if node.NodeType == "tunnel_client" {
config, err := service.GetFlaredTunnelConfig(node) config, err := service.GetFlaredTunnelConfig(node)
if err != nil { if err != nil {
respondFailure(c, "无法生成隧道配置: "+err.Error()) response.RespondFailure(c, "无法生成隧道配置: "+err.Error())
return return
} }
respondSuccess(c, config) response.RespondSuccess(c, config)
return return
} }
config, err := service.GetActiveConfigForAgent() config, err := service.GetActiveConfigForAgent()
if err != nil { if err != nil {
respondFailure(c, "当前没有激活版本") response.RespondFailure(c, "当前没有激活版本")
return return
} }
respondSuccess(c, config) response.RespondSuccess(c, config)
} }
// AgentReportApplyLog godoc // AgentReportApplyLog godoc
@@ -150,7 +152,7 @@ func AgentGetActiveConfig(c *gin.Context) {
// @Router /api/agent/apply-logs [post] // @Router /api/agent/apply-logs [post]
func AgentReportApplyLog(c *gin.Context) { func AgentReportApplyLog(c *gin.Context) {
var payload service.ApplyLogPayload var payload service.ApplyLogPayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
@@ -160,10 +162,10 @@ func AgentReportApplyLog(c *gin.Context) {
log, err := service.ReportApplyLog(payload) log, err := service.ReportApplyLog(payload)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, log) response.RespondSuccess(c, log)
} }
// AgentWebSocket godoc // AgentWebSocket godoc
@@ -174,7 +176,7 @@ func AgentReportApplyLog(c *gin.Context) {
func AgentWebSocket(c *gin.Context) { func AgentWebSocket(c *gin.Context) {
authNode, ok := c.Get("agent_node") authNode, ok := c.Get("agent_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作,Agent Token 无效") response.RespondUnauthorized(c, "无权进行此操作,Agent Token 无效")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
@@ -268,19 +270,19 @@ func handleAgentWSStatus(c *gin.Context, node *model.Node, message service.Agent
return return
} }
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
response, err := service.HeartbeatNode(freshNode, payload) res, err := service.HeartbeatNode(freshNode, payload)
if err != nil { if err != nil {
slog.Debug("agent ws status handling failed", "node_id", node.NodeID, "error", err) slog.Debug("agent ws status handling failed", "node_id", node.NodeID, "error", err)
return return
} }
settingsSent := service.SendAgentWSSettings(node.NodeID, response.AgentSettings) settingsSent := service.SendAgentWSSettings(node.NodeID, res.AgentSettings)
activeConfigSent := false activeConfigSent := false
if response.ActiveConfig != nil { if res.ActiveConfig != nil {
activeConfigSent = service.SendAgentWSActiveConfig(node.NodeID, response.ActiveConfig) activeConfigSent = service.SendAgentWSActiveConfig(node.NodeID, res.ActiveConfig)
} }
wafIPGroupsSent := false wafIPGroupsSent := false
if len(response.WAFIPGroups) > 0 { if len(res.WAFIPGroups) > 0 {
wafIPGroupsSent = service.SendAgentWSWAFIPGroups(node.NodeID, response.WAFIPGroups) wafIPGroupsSent = service.SendAgentWSWAFIPGroups(node.NodeID, res.WAFIPGroups)
} }
slog.Debug("agent ws status processed", slog.Debug("agent ws status processed",
"node_id", node.NodeID, "node_id", node.NodeID,
@@ -302,10 +304,10 @@ func handleAgentWSStatus(c *gin.Context, node *model.Node, message service.Agent
func GetNodes(c *gin.Context) { func GetNodes(c *gin.Context) {
nodes, err := service.ListNodeViews() nodes, err := service.ListNodeViews()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nodes) response.RespondSuccess(c, nodes)
} }
// GetApplyLogs godoc // GetApplyLogs godoc
@@ -323,10 +325,10 @@ func GetApplyLogs(c *gin.Context) {
PageSize: readIntQueryFallback(c, "pageSize", "page_size"), PageSize: readIntQueryFallback(c, "pageSize", "page_size"),
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, logs) response.RespondSuccess(c, logs)
} }
// CleanupApplyLogs godoc // CleanupApplyLogs godoc
@@ -339,15 +341,15 @@ func GetApplyLogs(c *gin.Context) {
// @Router /api/apply-logs/cleanup [post] // @Router /api/apply-logs/cleanup [post]
func CleanupApplyLogs(c *gin.Context) { func CleanupApplyLogs(c *gin.Context) {
var input service.ApplyLogCleanupInput var input service.ApplyLogCleanupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
result, err := service.CleanupApplyLogs(input) result, err := service.CleanupApplyLogs(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
func readIntQueryFallback(c *gin.Context, primary string, secondary string) int { func readIntQueryFallback(c *gin.Context, primary string, secondary string) int {
+54 -52
View File
@@ -5,6 +5,8 @@ import (
"fmt" "fmt"
"net/url" "net/url"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"strconv" "strconv"
@@ -49,139 +51,139 @@ func (payload authSourcePayload) toModel() model.AuthSource {
func ListAuthSources(c *gin.Context) { func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources() sources, err := model.GetAuthSources()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, sources) response.RespondSuccess(c, sources)
} }
func CreateAuthSource(c *gin.Context) { func CreateAuthSource(c *gin.Context) {
var payload authSourcePayload var payload authSourcePayload
if err := decodeJSONBody(c.Request.Body, &payload); err != nil { if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
source := payload.toModel() source := payload.toModel()
if err := model.CreateAuthSource(&source); err != nil { if err := model.CreateAuthSource(&source); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
source.Sanitize() source.Sanitize()
respondSuccess(c, source) response.RespondSuccess(c, source)
} }
func UpdateAuthSource(c *gin.Context) { func UpdateAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c) id, err := parseAuthSourceID(c)
if err != nil { if err != nil {
respondBadRequest(c, err.Error()) response.RespondBadRequest(c, err.Error())
return return
} }
var payload authSourcePayload var payload authSourcePayload
if err := decodeJSONBody(c.Request.Body, &payload); err != nil { if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
source := payload.toModel() source := payload.toModel()
source.ID = id source.ID = id
keepSecret := strings.TrimSpace(source.ClientSecret) == "" keepSecret := strings.TrimSpace(source.ClientSecret) == ""
if err := model.UpdateAuthSource(&source, keepSecret); err != nil { if err := model.UpdateAuthSource(&source, keepSecret); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
updated, err := model.GetAuthSourceByID(id) updated, err := model.GetAuthSourceByID(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
updated.Sanitize() updated.Sanitize()
respondSuccess(c, updated) response.RespondSuccess(c, updated)
} }
func DeleteAuthSource(c *gin.Context) { func DeleteAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c) id, err := parseAuthSourceID(c)
if err != nil { if err != nil {
respondBadRequest(c, err.Error()) response.RespondBadRequest(c, err.Error())
return return
} }
if err := model.DeleteAuthSource(id); err != nil { if err := model.DeleteAuthSource(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func ToggleAuthSource(c *gin.Context) { func ToggleAuthSource(c *gin.Context) {
id, err := parseAuthSourceID(c) id, err := parseAuthSourceID(c)
if err != nil { if err != nil {
respondBadRequest(c, err.Error()) response.RespondBadRequest(c, err.Error())
return return
} }
var payload authSourceTogglePayload var payload authSourceTogglePayload
if err := decodeJSONBody(c.Request.Body, &payload); err != nil { if err := bind.DecodeJSONBody(c.Request.Body, &payload); err != nil {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
if err := model.ToggleAuthSource(id, payload.IsActive); err != nil { if err := model.ToggleAuthSource(id, payload.IsActive); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func OAuthAuthorize(c *gin.Context) { func OAuthAuthorize(c *gin.Context) {
source, err := getAuthSourceFromRoute(c) source, err := getAuthSourceFromRoute(c)
if err != nil { if err != nil {
respondBadRequest(c, err.Error()) response.RespondBadRequest(c, err.Error())
return return
} }
if !source.IsActive { if !source.IsActive {
respondFailure(c, "认证源未启用") response.RespondFailure(c, "认证源未启用")
return return
} }
if err := source.Validate(); err != nil { if err := source.Validate(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
state, err := service.GenerateOAuthState() state, err := service.GenerateOAuthState()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
session := sessions.Default(c) session := sessions.Default(c)
session.Set(oauthStateSessionKey(source.ID), state) session.Set(oauthStateSessionKey(source.ID), state)
if err := session.Save(); err != nil { if err := session.Save(); err != nil {
respondFailure(c, "无法保存授权状态,请重试") response.RespondFailure(c, "无法保存授权状态,请重试")
return return
} }
redirectURL := oauthFrontendCallbackURL(c, source.ID) redirectURL := oauthFrontendCallbackURL(c, source.ID)
authorizeURL, err := service.BuildAuthorizeURL(c.Request.Context(), source, redirectURL, state) authorizeURL, err := service.BuildAuthorizeURL(c.Request.Context(), source, redirectURL, state)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, gin.H{"authorize_url": authorizeURL}) response.RespondSuccess(c, gin.H{"authorize_url": authorizeURL})
} }
func OAuthCallback(c *gin.Context) { func OAuthCallback(c *gin.Context) {
source, err := getAuthSourceFromRoute(c) source, err := getAuthSourceFromRoute(c)
if err != nil { if err != nil {
respondBadRequest(c, err.Error()) response.RespondBadRequest(c, err.Error())
return return
} }
if !source.IsActive { if !source.IsActive {
respondFailure(c, "认证源未启用") response.RespondFailure(c, "认证源未启用")
return return
} }
session := sessions.Default(c) session := sessions.Default(c)
expectedState, _ := session.Get(oauthStateSessionKey(source.ID)).(string) expectedState, _ := session.Get(oauthStateSessionKey(source.ID)).(string)
state := c.Query("state") state := c.Query("state")
if expectedState == "" || state == "" || state != expectedState { if expectedState == "" || state == "" || state != expectedState {
respondFailure(c, "授权状态无效,请重新登录") response.RespondFailure(c, "授权状态无效,请重新登录")
return return
} }
session.Delete(oauthStateSessionKey(source.ID)) session.Delete(oauthStateSessionKey(source.ID))
if err := session.Save(); err != nil { if err := session.Save(); err != nil {
respondFailure(c, "无法更新授权状态,请重试") response.RespondFailure(c, "无法更新授权状态,请重试")
return return
} }
if oauthError := c.Query("error"); oauthError != "" { if oauthError := c.Query("error"); oauthError != "" {
@@ -189,13 +191,13 @@ func OAuthCallback(c *gin.Context) {
if description == "" { if description == "" {
description = oauthError description = oauthError
} }
respondFailure(c, description) response.RespondFailure(c, description)
return return
} }
profile, err := service.ExchangeOAuthProfile(c.Request.Context(), source, c.Query("code"), oauthFrontendCallbackURL(c, source.ID)) profile, err := service.ExchangeOAuthProfile(c.Request.Context(), source, c.Query("code"), oauthFrontendCallbackURL(c, source.ID))
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
var currentUserID *int var currentUserID *int
@@ -204,91 +206,91 @@ func OAuthCallback(c *gin.Context) {
} }
result, pending, err := service.CompleteOAuthLogin(source, profile, currentUserID) result, pending, err := service.CompleteOAuthLogin(source, profile, currentUserID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
if pending != nil { if pending != nil {
raw, err := json.Marshal(pending) raw, err := json.Marshal(pending)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
session.Set(pendingExternalAccountSessionKey, string(raw)) session.Set(pendingExternalAccountSessionKey, string(raw))
if err := session.Save(); err != nil { if err := session.Save(); err != nil {
respondFailure(c, "无法保存待绑定账号,请重试") response.RespondFailure(c, "无法保存待绑定账号,请重试")
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
return return
} }
if result.User != nil { if result.User != nil {
cleanUser, err := setLoginToken(result.User) cleanUser, err := setLoginToken(result.User)
if err != nil { if err != nil {
respondFailure(c, "无法保存会话信息,请重试") response.RespondFailure(c, "无法保存会话信息,请重试")
return return
} }
result.User = cleanUser result.User = cleanUser
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
func LinkExistingOAuthAccount(c *gin.Context) { func LinkExistingOAuthAccount(c *gin.Context) {
session := sessions.Default(c) session := sessions.Default(c)
raw, _ := session.Get(pendingExternalAccountSessionKey).(string) raw, _ := session.Get(pendingExternalAccountSessionKey).(string)
if raw == "" { if raw == "" {
respondFailure(c, "待绑定第三方账号已失效,请重新登录") response.RespondFailure(c, "待绑定第三方账号已失效,请重新登录")
return return
} }
var pending service.PendingExternalAccount var pending service.PendingExternalAccount
if err := json.Unmarshal([]byte(raw), &pending); err != nil { if err := json.Unmarshal([]byte(raw), &pending); err != nil {
respondFailure(c, "待绑定第三方账号无效,请重新登录") response.RespondFailure(c, "待绑定第三方账号无效,请重新登录")
return return
} }
var input service.LinkExistingRequest var input service.LinkExistingRequest
if err := decodeJSONBody(c.Request.Body, &input); err != nil { if err := bind.DecodeJSONBody(c.Request.Body, &input); err != nil {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
user, err := service.LinkPendingExternalAccount(&pending, input) user, err := service.LinkPendingExternalAccount(&pending, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
session.Delete(pendingExternalAccountSessionKey) session.Delete(pendingExternalAccountSessionKey)
if err := session.Save(); err != nil { if err := session.Save(); err != nil {
respondFailure(c, "无法更新会话信息,请重试") response.RespondFailure(c, "无法更新会话信息,请重试")
return return
} }
cleanUser, err := setLoginToken(user) cleanUser, err := setLoginToken(user)
if err != nil { if err != nil {
respondFailure(c, "无法保存会话信息,请重试") response.RespondFailure(c, "无法保存会话信息,请重试")
return return
} }
respondSuccess(c, service.OAuthCallbackResult{Status: "linked", User: cleanUser}) response.RespondSuccess(c, service.OAuthCallbackResult{Status: "linked", User: cleanUser})
} }
func ListExternalAccounts(c *gin.Context) { func ListExternalAccounts(c *gin.Context) {
userID := c.GetInt("id") userID := c.GetInt("id")
accounts, err := model.ListExternalAccountsByUserID(userID) accounts, err := model.ListExternalAccountsByUserID(userID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, accounts) response.RespondSuccess(c, accounts)
} }
func DeleteExternalAccount(c *gin.Context) { func DeleteExternalAccount(c *gin.Context) {
rawID := strings.TrimSpace(c.Param("id")) rawID := strings.TrimSpace(c.Param("id"))
parsedID, err := strconv.ParseUint(rawID, 10, 64) parsedID, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || parsedID == 0 { if err != nil || parsedID == 0 {
respondBadRequest(c, "绑定记录 ID 无效") response.RespondBadRequest(c, "绑定记录 ID 无效")
return return
} }
if err := model.DeleteExternalAccountForUser(uint(parsedID), c.GetInt("id")); err != nil { if err := model.DeleteExternalAccountForUser(uint(parsedID), c.GetInt("id")); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func parseAuthSourceID(c *gin.Context) (uint, error) { func parseAuthSourceID(c *gin.Context) (uint, error) {
+49
View File
@@ -0,0 +1,49 @@
package bind
import (
"encoding/json"
"errors"
"io"
"strconv"
"github.com/gin-gonic/gin"
"openflare/common/response"
)
// DecodeJSONBody decodes JSON reader to target
func DecodeJSONBody(body io.Reader, target any) error {
return json.NewDecoder(body).Decode(target)
}
// OptionalJSON decodes optional JSON body of reader to target, allowing EOF
func OptionalJSON(body io.Reader, target any) error {
if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) {
return err
}
return nil
}
// IDParam parses "id" parameter from context path
func IDParam(c *gin.Context) (uint, bool) {
return IDParamByName(c, "id")
}
// IDParamByName parses target parameter from context path
func IDParamByName(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 {
response.RespondBadRequest(c, "")
return 0, false
}
return uint(id), true
}
// JSON binds JSON body of context request to target
func JSON(c *gin.Context, target any) bool {
if err := DecodeJSONBody(c.Request.Body, target); err != nil {
response.RespondBadRequest(c, "")
return false
}
return true
}
+21 -19
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -16,10 +18,10 @@ import (
func GetConfigVersions(c *gin.Context) { func GetConfigVersions(c *gin.Context) {
versions, err := service.ListConfigVersions() versions, err := service.ListConfigVersions()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, versions) response.RespondSuccess(c, versions)
} }
// GetConfigVersion godoc // GetConfigVersion godoc
@@ -32,16 +34,16 @@ func GetConfigVersions(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/config-versions/{id} [get] // @Router /api/config-versions/{id} [get]
func GetConfigVersion(c *gin.Context) { func GetConfigVersion(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
version, err := service.GetConfigVersionDetail(id) version, err := service.GetConfigVersionDetail(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, version) response.RespondSuccess(c, version)
} }
// GetActiveConfigVersion godoc // GetActiveConfigVersion godoc
@@ -54,10 +56,10 @@ func GetConfigVersion(c *gin.Context) {
func GetActiveConfigVersion(c *gin.Context) { func GetActiveConfigVersion(c *gin.Context) {
version, err := service.GetActiveConfigVersion() version, err := service.GetActiveConfigVersion()
if err != nil { if err != nil {
respondFailure(c, "当前没有激活版本") response.RespondFailure(c, "当前没有激活版本")
return return
} }
respondSuccess(c, version) response.RespondSuccess(c, version)
} }
// PreviewConfigVersion godoc // PreviewConfigVersion godoc
@@ -70,10 +72,10 @@ func GetActiveConfigVersion(c *gin.Context) {
func PreviewConfigVersion(c *gin.Context) { func PreviewConfigVersion(c *gin.Context) {
preview, err := service.PreviewConfigVersion() preview, err := service.PreviewConfigVersion()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, preview) response.RespondSuccess(c, preview)
} }
// DiffConfigVersion godoc // DiffConfigVersion godoc
@@ -86,10 +88,10 @@ func PreviewConfigVersion(c *gin.Context) {
func DiffConfigVersion(c *gin.Context) { func DiffConfigVersion(c *gin.Context) {
diff, err := service.DiffConfigVersion() diff, err := service.DiffConfigVersion()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, diff) response.RespondSuccess(c, diff)
} }
// PublishConfigVersion godoc // PublishConfigVersion godoc
@@ -104,10 +106,10 @@ func PublishConfigVersion(c *gin.Context) {
force := c.Query("force") == "true" force := c.Query("force") == "true"
result, err := service.PublishConfigVersion(username, force) result, err := service.PublishConfigVersion(username, force)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result.Version) response.RespondSuccess(c, result.Version)
} }
// ActivateConfigVersion godoc // ActivateConfigVersion godoc
@@ -120,16 +122,16 @@ func PublishConfigVersion(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/config-versions/{id}/activate [post] // @Router /api/config-versions/{id}/activate [post]
func ActivateConfigVersion(c *gin.Context) { func ActivateConfigVersion(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
version, err := service.ActivateConfigVersion(id) version, err := service.ActivateConfigVersion(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, version) response.RespondSuccess(c, version)
} }
type CleanupConfigVersionRequest struct { type CleanupConfigVersionRequest struct {
@@ -147,17 +149,17 @@ type CleanupConfigVersionRequest struct {
// @Router /api/config-versions/cleanup [post] // @Router /api/config-versions/cleanup [post]
func CleanupConfigVersions(c *gin.Context) { func CleanupConfigVersions(c *gin.Context) {
var req CleanupConfigVersionRequest var req CleanupConfigVersionRequest
if !bindJSON(c, &req) { if !bind.JSON(c, &req) {
return return
} }
deletedCount, err := service.CleanupConfigVersions(req.KeepCount) deletedCount, err := service.CleanupConfigVersions(req.KeepCount)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessWithExtras(c, map[string]interface{}{"deleted_count": deletedCount}, gin.H{ response.RespondSuccessWithExtras(c, map[string]interface{}{"deleted_count": deletedCount}, gin.H{
"message": "清理成功", "message": "清理成功",
}) })
} }
+3 -2
View File
@@ -1,6 +1,7 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -40,10 +41,10 @@ type dashboardTrendsPayload struct {
func GetDashboardOverview(c *gin.Context) { func GetDashboardOverview(c *gin.Context) {
view, err := service.GetDashboardOverview() view, err := service.GetDashboardOverview()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, compressDashboardOverview(view)) response.RespondSuccess(c, compressDashboardOverview(view))
} }
func compressDashboardOverview(view *service.DashboardOverviewView) *dashboardOverviewPayload { func compressDashboardOverview(view *service.DashboardOverviewView) *dashboardOverviewPayload {
+6 -4
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -16,14 +18,14 @@ import (
// @Router /api/option/database/cleanup [post] // @Router /api/option/database/cleanup [post]
func CleanupDatabaseObservability(c *gin.Context) { func CleanupDatabaseObservability(c *gin.Context) {
var input service.DatabaseCleanupInput var input service.DatabaseCleanupInput
if err := decodeOptionalJSONBody(c.Request.Body, &input); err != nil { if err := bind.OptionalJSON(c.Request.Body, &input); err != nil {
respondBadRequest(c, "") response.RespondBadRequest(c, "")
return return
} }
result, err := service.CleanupDatabaseObservability(input) result, err := service.CleanupDatabaseObservability(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
+17 -15
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -22,10 +24,10 @@ type DnsAccountInput struct {
func GetDnsAccounts(c *gin.Context) { func GetDnsAccounts(c *gin.Context) {
accounts, err := model.ListDnsAccounts() accounts, err := model.ListDnsAccounts()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, accounts) response.RespondSuccess(c, accounts)
} }
// CreateDnsAccount godoc // CreateDnsAccount godoc
@@ -39,7 +41,7 @@ func GetDnsAccounts(c *gin.Context) {
// @Router /api/dns-accounts/ [post] // @Router /api/dns-accounts/ [post]
func CreateDnsAccount(c *gin.Context) { func CreateDnsAccount(c *gin.Context) {
var input DnsAccountInput var input DnsAccountInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
@@ -50,11 +52,11 @@ func CreateDnsAccount(c *gin.Context) {
} }
if err := account.Insert(); err != nil { if err := account.Insert(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, account) response.RespondSuccess(c, account)
} }
// UpdateDnsAccount godoc // UpdateDnsAccount godoc
@@ -68,19 +70,19 @@ func CreateDnsAccount(c *gin.Context) {
// @Success 200 {object} map[string]interface{} // @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/{id}/update [post] // @Router /api/dns-accounts/{id}/update [post]
func UpdateDnsAccount(c *gin.Context) { func UpdateDnsAccount(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input DnsAccountInput var input DnsAccountInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
account, err := model.GetDnsAccountByID(id) account, err := model.GetDnsAccountByID(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
@@ -89,11 +91,11 @@ func UpdateDnsAccount(c *gin.Context) {
account.Authorization = input.Authorization account.Authorization = input.Authorization
if err := account.Update(); err != nil { if err := account.Update(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, account) response.RespondSuccess(c, account)
} }
// DeleteDnsAccount godoc // DeleteDnsAccount godoc
@@ -105,14 +107,14 @@ func UpdateDnsAccount(c *gin.Context) {
// @Success 200 {object} map[string]interface{} // @Success 200 {object} map[string]interface{}
// @Router /api/dns-accounts/{id}/delete [post] // @Router /api/dns-accounts/{id}/delete [post]
func DeleteDnsAccount(c *gin.Context) { func DeleteDnsAccount(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
account, err := model.GetDnsAccountByID(id) account, err := model.GetDnsAccountByID(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
@@ -120,14 +122,14 @@ func DeleteDnsAccount(c *gin.Context) {
var count int64 var count int64
model.DB.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", id).Count(&count) model.DB.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", id).Count(&count)
if count > 0 { if count > 0 {
respondFailure(c, "该 DNS 账号已被证书使用,无法删除") response.RespondFailure(c, "该 DNS 账号已被证书使用,无法删除")
return return
} }
if err := account.Delete(); err != nil { if err := account.Delete(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
+14 -12
View File
@@ -4,6 +4,8 @@ import (
"log/slog" "log/slog"
"net" "net"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"time" "time"
@@ -24,21 +26,21 @@ import (
// @Router /api/flared/heartbeat [post] // @Router /api/flared/heartbeat [post]
func FlaredHeartbeat(c *gin.Context) { func FlaredHeartbeat(c *gin.Context) {
var payload service.FlaredHeartbeatPayload var payload service.FlaredHeartbeatPayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
authNode, ok := c.Get("flared_node") authNode, ok := c.Get("flared_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作,Tunnel Token 无效") response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
response, err := service.HeartbeatFlared(node, payload) res, err := service.HeartbeatFlared(node, payload)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, response) response.RespondSuccess(c, res)
} }
// FlaredGetActiveConfig godoc // FlaredGetActiveConfig godoc
@@ -51,16 +53,16 @@ func FlaredHeartbeat(c *gin.Context) {
func FlaredGetActiveConfig(c *gin.Context) { func FlaredGetActiveConfig(c *gin.Context) {
authNode, ok := c.Get("flared_node") authNode, ok := c.Get("flared_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作,Tunnel Token 无效") response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
config, err := service.GetFlaredTunnelConfig(node) config, err := service.GetFlaredTunnelConfig(node)
if err != nil { if err != nil {
respondFailure(c, "无法生成隧道配置: "+err.Error()) response.RespondFailure(c, "无法生成隧道配置: "+err.Error())
return return
} }
respondSuccess(c, config) response.RespondSuccess(c, config)
} }
// FlaredReportApplyLog godoc // FlaredReportApplyLog godoc
@@ -74,7 +76,7 @@ func FlaredGetActiveConfig(c *gin.Context) {
// @Router /api/flared/apply-log [post] // @Router /api/flared/apply-log [post]
func FlaredReportApplyLog(c *gin.Context) { func FlaredReportApplyLog(c *gin.Context) {
var payload service.ApplyLogPayload var payload service.ApplyLogPayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
if authNode, ok := c.Get("flared_node"); ok { if authNode, ok := c.Get("flared_node"); ok {
@@ -82,10 +84,10 @@ func FlaredReportApplyLog(c *gin.Context) {
} }
log, err := service.ReportApplyLog(payload) log, err := service.ReportApplyLog(payload)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, log) response.RespondSuccess(c, log)
} }
// FlaredWebSocket godoc // FlaredWebSocket godoc
@@ -96,7 +98,7 @@ func FlaredReportApplyLog(c *gin.Context) {
func FlaredWebSocket(c *gin.Context) { func FlaredWebSocket(c *gin.Context) {
authNode, ok := c.Get("flared_node") authNode, ok := c.Get("flared_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作,Tunnel Token 无效") response.RespondUnauthorized(c, "无权进行此操作,Tunnel Token 无效")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
+5 -3
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -22,14 +24,14 @@ type geoIPLookupRequest struct {
// @Router /api/option/geoip/lookup [post] // @Router /api/option/geoip/lookup [post]
func LookupGeoIP(c *gin.Context) { func LookupGeoIP(c *gin.Context) {
var request geoIPLookupRequest var request geoIPLookupRequest
if !bindJSON(c, &request) { if !bind.JSON(c, &request) {
return return
} }
view, err := service.LookupGeoIP(request.Provider, request.IP) view, err := service.LookupGeoIP(request.Provider, request.IP)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, view) response.RespondSuccess(c, view)
} }
+13 -12
View File
@@ -8,6 +8,7 @@ import (
"log/slog" "log/slog"
"net/http" "net/http"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/model" "openflare/model"
"time" "time"
@@ -84,13 +85,13 @@ func GitHubOAuth(c *gin.Context) {
} }
if !common.GitHubOAuthEnabled { if !common.GitHubOAuthEnabled {
respondFailure(c, "管理员未开启通过 GitHub 登录以及注册") response.RespondFailure(c, "管理员未开启通过 GitHub 登录以及注册")
return return
} }
code := c.Query("code") code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code) githubUser, err := getGitHubUserInfoByCode(code)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user := model.User{ user := model.User{
@@ -99,16 +100,16 @@ func GitHubOAuth(c *gin.Context) {
if model.IsGitHubIdAlreadyTaken(user.GitHubId) { if model.IsGitHubIdAlreadyTaken(user.GitHubId) {
err := user.FillUserByGitHubId() err := user.FillUserByGitHubId()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
} else { } else {
respondFailure(c, "管理员关闭了新用户注册") response.RespondFailure(c, "管理员关闭了新用户注册")
return return
} }
if user.Status != common.UserStatusEnabled { if user.Status != common.UserStatusEnabled {
respondFailure(c, "用户已被封禁") response.RespondFailure(c, "用户已被封禁")
return return
} }
setupLogin(&user, c) setupLogin(&user, c)
@@ -116,38 +117,38 @@ func GitHubOAuth(c *gin.Context) {
func GitHubBind(c *gin.Context) { func GitHubBind(c *gin.Context) {
if !common.GitHubOAuthEnabled { if !common.GitHubOAuthEnabled {
respondFailure(c, "管理员未开启通过 GitHub 登录以及注册") response.RespondFailure(c, "管理员未开启通过 GitHub 登录以及注册")
return return
} }
code := c.Query("code") code := c.Query("code")
githubUser, err := getGitHubUserInfoByCode(code) githubUser, err := getGitHubUserInfoByCode(code)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user := model.User{ user := model.User{
GitHubId: githubUser.Login, GitHubId: githubUser.Login,
} }
if model.IsGitHubIdAlreadyTaken(user.GitHubId) { if model.IsGitHubIdAlreadyTaken(user.GitHubId) {
respondFailure(c, "该 GitHub 账户已被绑定") response.RespondFailure(c, "该 GitHub 账户已被绑定")
return return
} }
currentUser := currentUserFromOpenFlareToken(c) currentUser := currentUserFromOpenFlareToken(c)
if currentUser == nil { if currentUser == nil {
respondFailure(c, "无权进行此操作,未登录或 token 无效") response.RespondFailure(c, "无权进行此操作,未登录或 token 无效")
return return
} }
user.Id = currentUser.Id user.Id = currentUser.Id
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user.GitHubId = githubUser.Login user.GitHubId = githubUser.Login
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "bind") response.RespondSuccessMessage(c, "bind")
} }
+16 -14
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"strings" "strings"
@@ -17,10 +19,10 @@ import (
func GetManagedDomains(c *gin.Context) { func GetManagedDomains(c *gin.Context) {
domains, err := service.ListManagedDomains() domains, err := service.ListManagedDomains()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, domains) response.RespondSuccess(c, domains)
} }
// CreateManagedDomain godoc // CreateManagedDomain godoc
@@ -35,15 +37,15 @@ func GetManagedDomains(c *gin.Context) {
// @Router /api/managed-domains/ [post] // @Router /api/managed-domains/ [post]
func CreateManagedDomain(c *gin.Context) { func CreateManagedDomain(c *gin.Context) {
var input service.ManagedDomainInput var input service.ManagedDomainInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
domain, err := service.CreateManagedDomain(input) domain, err := service.CreateManagedDomain(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, domain) response.RespondSuccess(c, domain)
} }
// UpdateManagedDomain godoc // UpdateManagedDomain godoc
@@ -58,20 +60,20 @@ func CreateManagedDomain(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/managed-domains/{id}/update [post] // @Router /api/managed-domains/{id}/update [post]
func UpdateManagedDomain(c *gin.Context) { func UpdateManagedDomain(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.ManagedDomainInput var input service.ManagedDomainInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
domain, err := service.UpdateManagedDomain(id, input) domain, err := service.UpdateManagedDomain(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, domain) response.RespondSuccess(c, domain)
} }
// DeleteManagedDomain godoc // DeleteManagedDomain godoc
@@ -84,15 +86,15 @@ func UpdateManagedDomain(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/managed-domains/{id}/delete [post] // @Router /api/managed-domains/{id}/delete [post]
func DeleteManagedDomain(c *gin.Context) { func DeleteManagedDomain(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteManagedDomain(id); err != nil { if err := service.DeleteManagedDomain(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
// MatchManagedDomainCertificate godoc // MatchManagedDomainCertificate godoc
@@ -107,8 +109,8 @@ func MatchManagedDomainCertificate(c *gin.Context) {
domain := strings.TrimSpace(c.Query("domain")) domain := strings.TrimSpace(c.Query("domain"))
result, err := service.MatchManagedDomainCertificate(domain) result, err := service.MatchManagedDomainCertificate(domain)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
+18 -16
View File
@@ -3,6 +3,8 @@ package controller
import ( import (
"fmt" "fmt"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"openflare/utils/mail" "openflare/utils/mail"
@@ -23,7 +25,7 @@ func GetStatus(c *gin.Context) {
if err != nil { if err != nil {
authSources = []service.PublicAuthSource{} authSources = []service.PublicAuthSource{}
} }
respondSuccess(c, gin.H{ response.RespondSuccess(c, gin.H{
"version": common.Version, "version": common.Version,
"start_time": common.StartTime, "start_time": common.StartTime,
"email_verification": common.EmailVerificationEnabled, "email_verification": common.EmailVerificationEnabled,
@@ -43,23 +45,23 @@ func GetStatus(c *gin.Context) {
func GetNotice(c *gin.Context) { func GetNotice(c *gin.Context) {
common.OptionMapRWMutex.RLock() common.OptionMapRWMutex.RLock()
defer common.OptionMapRWMutex.RUnlock() defer common.OptionMapRWMutex.RUnlock()
respondSuccess(c, common.OptionMap["Notice"]) response.RespondSuccess(c, common.OptionMap["Notice"])
} }
func GetAbout(c *gin.Context) { func GetAbout(c *gin.Context) {
common.OptionMapRWMutex.RLock() common.OptionMapRWMutex.RLock()
defer common.OptionMapRWMutex.RUnlock() defer common.OptionMapRWMutex.RUnlock()
respondSuccess(c, common.OptionMap["About"]) response.RespondSuccess(c, common.OptionMap["About"])
} }
func SendEmailVerification(c *gin.Context) { func SendEmailVerification(c *gin.Context) {
email := c.Query("email") email := c.Query("email")
if err := validation.Validate.Var(email, "required,email"); err != nil { if err := validation.Validate.Var(email, "required,email"); err != nil {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
if model.IsEmailAlreadyTaken(email) { if model.IsEmailAlreadyTaken(email) {
respondFailure(c, "邮箱地址已被占用") response.RespondFailure(c, "邮箱地址已被占用")
return return
} }
code := security.GenerateVerificationCode(6) code := security.GenerateVerificationCode(6)
@@ -77,20 +79,20 @@ func SendEmailVerification(c *gin.Context) {
} }
err := mail.SendEmail(cfg, subject, email, content) err := mail.SendEmail(cfg, subject, email, content)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func SendPasswordResetEmail(c *gin.Context) { func SendPasswordResetEmail(c *gin.Context) {
email := c.Query("email") email := c.Query("email")
if err := validation.Validate.Var(email, "required,email"); err != nil { if err := validation.Validate.Var(email, "required,email"); err != nil {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
if !model.IsEmailAlreadyTaken(email) { if !model.IsEmailAlreadyTaken(email) {
respondFailure(c, "该邮箱地址未注册") response.RespondFailure(c, "该邮箱地址未注册")
return return
} }
code := security.GenerateVerificationCode(0) code := security.GenerateVerificationCode(0)
@@ -109,10 +111,10 @@ func SendPasswordResetEmail(c *gin.Context) {
} }
err := mail.SendEmail(cfg, subject, email, content) err := mail.SendEmail(cfg, subject, email, content)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
type PasswordResetRequest struct { type PasswordResetRequest struct {
@@ -122,23 +124,23 @@ type PasswordResetRequest struct {
func ResetPassword(c *gin.Context) { func ResetPassword(c *gin.Context) {
var req PasswordResetRequest var req PasswordResetRequest
if !bindJSON(c, &req) { if !bind.JSON(c, &req) {
return return
} }
if req.Email == "" || req.Token == "" { if req.Email == "" || req.Token == "" {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
if !security.VerifyCodeWithKey(req.Email, req.Token, security.PasswordResetPurpose) { if !security.VerifyCodeWithKey(req.Email, req.Token, security.PasswordResetPurpose) {
respondFailure(c, "重置链接非法或已过期") response.RespondFailure(c, "重置链接非法或已过期")
return return
} }
password := security.GenerateVerificationCode(12) password := security.GenerateVerificationCode(12)
err := model.ResetUserPasswordByEmail(req.Email, password) err := model.ResetUserPasswordByEmail(req.Email, password)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
security.DeleteKey(req.Email, security.PasswordResetPurpose) security.DeleteKey(req.Email, security.PasswordResetPurpose)
respondSuccess(c, password) response.RespondSuccess(c, password)
} }
+37 -35
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -28,16 +30,16 @@ type nodeObservabilityQuery struct {
// @Router /api/nodes/ [post] // @Router /api/nodes/ [post]
func CreateNode(c *gin.Context) { func CreateNode(c *gin.Context) {
var input service.NodeInput var input service.NodeInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
node, err := service.CreateNode(input) node, err := service.CreateNode(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, node) response.RespondSuccess(c, node)
} }
// GetNodeBootstrapToken godoc // GetNodeBootstrapToken godoc
@@ -50,10 +52,10 @@ func CreateNode(c *gin.Context) {
func GetNodeBootstrapToken(c *gin.Context) { func GetNodeBootstrapToken(c *gin.Context) {
bootstrap, err := service.GetNodeBootstrapView() bootstrap, err := service.GetNodeBootstrapView()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, bootstrap) response.RespondSuccess(c, bootstrap)
} }
// RotateNodeBootstrapToken godoc // RotateNodeBootstrapToken godoc
@@ -66,10 +68,10 @@ func GetNodeBootstrapToken(c *gin.Context) {
func RotateNodeBootstrapToken(c *gin.Context) { func RotateNodeBootstrapToken(c *gin.Context) {
bootstrap, err := service.RotateGlobalDiscoveryToken() bootstrap, err := service.RotateGlobalDiscoveryToken()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, bootstrap) response.RespondSuccess(c, bootstrap)
} }
// UpdateNode godoc // UpdateNode godoc
@@ -84,22 +86,22 @@ func RotateNodeBootstrapToken(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/update [post] // @Router /api/nodes/{id}/update [post]
func UpdateNode(c *gin.Context) { func UpdateNode(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.NodeInput var input service.NodeInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
node, err := service.UpdateNode(id, input) node, err := service.UpdateNode(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, node) response.RespondSuccess(c, node)
} }
// DeleteNode godoc // DeleteNode godoc
@@ -112,16 +114,16 @@ func UpdateNode(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/delete [post] // @Router /api/nodes/{id}/delete [post]
func DeleteNode(c *gin.Context) { func DeleteNode(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteNode(id); err != nil { if err := service.DeleteNode(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
// RequestNodeAgentUpdate godoc // RequestNodeAgentUpdate godoc
@@ -134,15 +136,15 @@ func DeleteNode(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/agent-update [post] // @Router /api/nodes/{id}/agent-update [post]
func RequestNodeAgentUpdate(c *gin.Context) { func RequestNodeAgentUpdate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var request nodeAgentUpdateRequest var request nodeAgentUpdateRequest
if c.Request.ContentLength > 0 { if c.Request.ContentLength > 0 {
if err := decodeOptionalJSONBody(c.Request.Body, &request); err != nil { if err := bind.OptionalJSON(c.Request.Body, &request); err != nil {
respondBadRequest(c, "") response.RespondBadRequest(c, "")
return return
} }
} }
@@ -152,10 +154,10 @@ func RequestNodeAgentUpdate(c *gin.Context) {
TagName: request.TagName, TagName: request.TagName,
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, node) response.RespondSuccess(c, node)
} }
// RequestNodeOpenrestyRestart godoc // RequestNodeOpenrestyRestart godoc
@@ -168,17 +170,17 @@ func RequestNodeAgentUpdate(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/openresty-restart [post] // @Router /api/nodes/{id}/openresty-restart [post]
func RequestNodeOpenrestyRestart(c *gin.Context) { func RequestNodeOpenrestyRestart(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
node, err := service.RequestNodeOpenrestyRestart(id) node, err := service.RequestNodeOpenrestyRestart(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, node) response.RespondSuccess(c, node)
} }
// RequestNodeForceSync godoc // RequestNodeForceSync godoc
@@ -191,17 +193,17 @@ func RequestNodeOpenrestyRestart(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/force-sync [post] // @Router /api/nodes/{id}/force-sync [post]
func RequestNodeForceSync(c *gin.Context) { func RequestNodeForceSync(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
node, err := service.RequestNodeForceSync(id) node, err := service.RequestNodeForceSync(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, node) response.RespondSuccess(c, node)
} }
// GetNodeAgentRelease godoc // GetNodeAgentRelease godoc
@@ -215,17 +217,17 @@ func RequestNodeForceSync(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/agent-release [get] // @Router /api/nodes/{id}/agent-release [get]
func GetNodeAgentRelease(c *gin.Context) { func GetNodeAgentRelease(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
release, err := service.GetNodeAgentRelease(c.Request.Context(), id, c.Query("channel")) release, err := service.GetNodeAgentRelease(c.Request.Context(), id, c.Query("channel"))
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, release) response.RespondSuccess(c, release)
} }
// GetNodeObservability godoc // GetNodeObservability godoc
@@ -240,14 +242,14 @@ func GetNodeAgentRelease(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/observability [get] // @Router /api/nodes/{id}/observability [get]
func GetNodeObservability(c *gin.Context) { func GetNodeObservability(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var query nodeObservabilityQuery var query nodeObservabilityQuery
if err := c.ShouldBindQuery(&query); err != nil { if err := c.ShouldBindQuery(&query); err != nil {
respondBadRequest(c, "") response.RespondBadRequest(c, "")
return return
} }
@@ -256,10 +258,10 @@ func GetNodeObservability(c *gin.Context) {
Limit: query.Limit, Limit: query.Limit,
}) })
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, view) response.RespondSuccess(c, view)
} }
// CleanupNodeHealthEvents godoc // CleanupNodeHealthEvents godoc
@@ -272,15 +274,15 @@ func GetNodeObservability(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/nodes/{id}/observability/cleanup [post] // @Router /api/nodes/{id}/observability/cleanup [post]
func CleanupNodeHealthEvents(c *gin.Context) { func CleanupNodeHealthEvents(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
result, err := service.CleanupNodeHealthEvents(id) result, err := service.CleanupNodeHealthEvents(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
+11 -9
View File
@@ -3,6 +3,8 @@ package controller
import ( import (
"fmt" "fmt"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"openflare/utils" "openflare/utils"
@@ -356,7 +358,7 @@ func GetOptions(c *gin.Context) {
}) })
} }
common.OptionMapRWMutex.RUnlock() common.OptionMapRWMutex.RUnlock()
respondSuccess(c, options) response.RespondSuccess(c, options)
} }
// UpdateOption godoc // UpdateOption godoc
@@ -370,20 +372,20 @@ func GetOptions(c *gin.Context) {
// @Router /api/option/update [post] // @Router /api/option/update [post]
func UpdateOption(c *gin.Context) { func UpdateOption(c *gin.Context) {
var option model.Option var option model.Option
if !bindJSON(c, &option) { if !bind.JSON(c, &option) {
return return
} }
state := buildOptionValidationState([]model.Option{option}) state := buildOptionValidationState([]model.Option{option})
if err := validateOptionWithState(option, state); err != nil { if err := validateOptionWithState(option, state); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
err := model.UpdateOption(option.Key, option.Value) err := model.UpdateOption(option.Key, option.Value)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
// UpdateOptionsBatch godoc // UpdateOptionsBatch godoc
@@ -397,18 +399,18 @@ func UpdateOption(c *gin.Context) {
// @Router /api/option/update-batch [post] // @Router /api/option/update-batch [post]
func UpdateOptionsBatch(c *gin.Context) { func UpdateOptionsBatch(c *gin.Context) {
var payload optionBatchPayload var payload optionBatchPayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
if len(payload.Options) == 0 { if len(payload.Options) == 0 {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
if err := updateOptions(payload.Options); err != nil { if err := updateOptions(payload.Options); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
+17 -15
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -9,63 +11,63 @@ import (
func GetOrigins(c *gin.Context) { func GetOrigins(c *gin.Context) {
origins, err := service.ListOrigins() origins, err := service.ListOrigins()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, origins) response.RespondSuccess(c, origins)
} }
func GetOrigin(c *gin.Context) { func GetOrigin(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
origin, err := service.GetOriginDetail(id) origin, err := service.GetOriginDetail(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, origin) response.RespondSuccess(c, origin)
} }
func CreateOrigin(c *gin.Context) { func CreateOrigin(c *gin.Context) {
var input service.OriginInput var input service.OriginInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
origin, err := service.CreateOrigin(input) origin, err := service.CreateOrigin(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, origin) response.RespondSuccess(c, origin)
} }
func UpdateOrigin(c *gin.Context) { func UpdateOrigin(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.OriginInput var input service.OriginInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
origin, err := service.UpdateOrigin(id, input) origin, err := service.UpdateOrigin(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, origin) response.RespondSuccess(c, origin)
} }
func DeleteOrigin(c *gin.Context) { func DeleteOrigin(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteOrigin(id); err != nil { if err := service.DeleteOrigin(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
+37 -35
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -9,88 +11,88 @@ import (
func ListPagesProjects(c *gin.Context) { func ListPagesProjects(c *gin.Context) {
projects, err := service.ListPagesProjects() projects, err := service.ListPagesProjects()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, projects) response.RespondSuccess(c, projects)
} }
func GetPagesProject(c *gin.Context) { func GetPagesProject(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
project, err := service.GetPagesProject(id) project, err := service.GetPagesProject(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, project) response.RespondSuccess(c, project)
} }
func CreatePagesProject(c *gin.Context) { func CreatePagesProject(c *gin.Context) {
var input service.PagesProjectInput var input service.PagesProjectInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
project, err := service.CreatePagesProject(input) project, err := service.CreatePagesProject(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, project) response.RespondSuccess(c, project)
} }
func UpdatePagesProject(c *gin.Context) { func UpdatePagesProject(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.PagesProjectInput var input service.PagesProjectInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
project, err := service.UpdatePagesProject(id, input) project, err := service.UpdatePagesProject(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, project) response.RespondSuccess(c, project)
} }
func DeletePagesProject(c *gin.Context) { func DeletePagesProject(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeletePagesProject(id); err != nil { if err := service.DeletePagesProject(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
func ListPagesDeployments(c *gin.Context) { func ListPagesDeployments(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
deployments, err := service.ListPagesProjectDeployments(id) deployments, err := service.ListPagesProjectDeployments(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, deployments) response.RespondSuccess(c, deployments)
} }
func UploadPagesDeployment(c *gin.Context) { func UploadPagesDeployment(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
file, err := c.FormFile("package") file, err := c.FormFile("package")
if err != nil { if err != nil {
respondBadRequest(c, "缺少 Pages 部署包") response.RespondBadRequest(c, "缺少 Pages 部署包")
return return
} }
deployment, err := service.UploadPagesDeployment( deployment, err := service.UploadPagesDeployment(
@@ -101,66 +103,66 @@ func UploadPagesDeployment(c *gin.Context) {
c.GetString("username"), c.GetString("username"),
) )
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, deployment) response.RespondSuccess(c, deployment)
} }
func ActivatePagesDeployment(c *gin.Context) { func ActivatePagesDeployment(c *gin.Context) {
projectID, ok := parseIDParam(c) projectID, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
deploymentID, ok := parseIDParamByName(c, "deployment_id") deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok { if !ok {
return return
} }
project, err := service.ActivatePagesDeployment(projectID, deploymentID) project, err := service.ActivatePagesDeployment(projectID, deploymentID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, project) response.RespondSuccess(c, project)
} }
func DeletePagesDeployment(c *gin.Context) { func DeletePagesDeployment(c *gin.Context) {
projectID, ok := parseIDParam(c) projectID, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
deploymentID, ok := parseIDParamByName(c, "deployment_id") deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok { if !ok {
return return
} }
if err := service.DeletePagesDeployment(projectID, deploymentID); err != nil { if err := service.DeletePagesDeployment(projectID, deploymentID); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
func ListPagesDeploymentFiles(c *gin.Context) { func ListPagesDeploymentFiles(c *gin.Context) {
deploymentID, ok := parseIDParamByName(c, "deployment_id") deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok { if !ok {
return return
} }
files, err := service.ListPagesDeploymentFiles(deploymentID) files, err := service.ListPagesDeploymentFiles(deploymentID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, files) response.RespondSuccess(c, files)
} }
func AgentDownloadPagesDeploymentPackage(c *gin.Context) { func AgentDownloadPagesDeploymentPackage(c *gin.Context) {
deploymentID, ok := parseIDParamByName(c, "deployment_id") deploymentID, ok := bind.IDParamByName(c, "deployment_id")
if !ok { if !ok {
return return
} }
filePath, fileName, err := service.GetPagesDeploymentPackagePath(deploymentID) filePath, fileName, err := service.GetPagesDeploymentPackagePath(deploymentID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
c.Header("Content-Disposition", "attachment; filename="+fileName) c.Header("Content-Disposition", "attachment; filename="+fileName)
+17 -15
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -16,10 +18,10 @@ import (
func GetProxyRoutes(c *gin.Context) { func GetProxyRoutes(c *gin.Context) {
routes, err := service.ListProxyRoutes() routes, err := service.ListProxyRoutes()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, routes) response.RespondSuccess(c, routes)
} }
// GetProxyRoute godoc // GetProxyRoute godoc
@@ -32,16 +34,16 @@ func GetProxyRoutes(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id} [get] // @Router /api/proxy-routes/{id} [get]
func GetProxyRoute(c *gin.Context) { func GetProxyRoute(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
route, err := service.GetProxyRoute(id) route, err := service.GetProxyRoute(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, route) response.RespondSuccess(c, route)
} }
// CreateProxyRoute godoc // CreateProxyRoute godoc
@@ -56,15 +58,15 @@ func GetProxyRoute(c *gin.Context) {
// @Router /api/proxy-routes/ [post] // @Router /api/proxy-routes/ [post]
func CreateProxyRoute(c *gin.Context) { func CreateProxyRoute(c *gin.Context) {
var input service.ProxyRouteInput var input service.ProxyRouteInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
route, err := service.CreateProxyRoute(input) route, err := service.CreateProxyRoute(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, route) response.RespondSuccess(c, route)
} }
// UpdateProxyRoute godoc // UpdateProxyRoute godoc
@@ -79,20 +81,20 @@ func CreateProxyRoute(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id}/update [post] // @Router /api/proxy-routes/{id}/update [post]
func UpdateProxyRoute(c *gin.Context) { func UpdateProxyRoute(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.ProxyRouteInput var input service.ProxyRouteInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
route, err := service.UpdateProxyRoute(id, input) route, err := service.UpdateProxyRoute(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, route) response.RespondSuccess(c, route)
} }
// DeleteProxyRoute godoc // DeleteProxyRoute godoc
@@ -105,13 +107,13 @@ func UpdateProxyRoute(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/proxy-routes/{id}/delete [post] // @Router /api/proxy-routes/{id}/delete [post]
func DeleteProxyRoute(c *gin.Context) { func DeleteProxyRoute(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteProxyRoute(id); err != nil { if err := service.DeleteProxyRoute(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
+7 -5
View File
@@ -3,6 +3,8 @@ package controller
import ( import (
"log/slog" "log/slog"
"net" "net"
"openflare/common/response"
"openflare/controller/bind"
"openflare/model" "openflare/model"
"openflare/service" "openflare/service"
"time" "time"
@@ -23,22 +25,22 @@ import (
// @Router /api/relay/heartbeat [post] // @Router /api/relay/heartbeat [post]
func RelayHeartbeat(c *gin.Context) { func RelayHeartbeat(c *gin.Context) {
var payload service.RelayHeartbeatPayload var payload service.RelayHeartbeatPayload
if !bindJSON(c, &payload) { if !bind.JSON(c, &payload) {
return return
} }
payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr) payload.IP = service.ResolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := c.Get("relay_node") authNode, ok := c.Get("relay_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作") response.RespondUnauthorized(c, "无权进行此操作")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
result, err := service.HeartbeatRelay(node, payload) result, err := service.HeartbeatRelay(node, payload)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
// RelayWebSocket godoc // RelayWebSocket godoc
@@ -49,7 +51,7 @@ func RelayHeartbeat(c *gin.Context) {
func RelayWebSocket(c *gin.Context) { func RelayWebSocket(c *gin.Context) {
authNode, ok := c.Get("relay_node") authNode, ok := c.Get("relay_node")
if !ok { if !ok {
respondUnauthorized(c, "无权进行此操作") response.RespondUnauthorized(c, "无权进行此操作")
return return
} }
node := authNode.(*model.Node) node := authNode.(*model.Node)
-68
View File
@@ -1,68 +0,0 @@
package controller
import (
"encoding/json"
"errors"
"io"
"strconv"
"github.com/gin-gonic/gin"
"openflare/common/response"
)
func respondSuccess(c *gin.Context, data any) {
response.RespondSuccess(c, data)
}
func respondSuccessWithExtras(c *gin.Context, data any, extras gin.H) {
response.RespondSuccessWithExtras(c, data, extras)
}
func respondSuccessMessage(c *gin.Context, message string) {
response.RespondSuccessMessage(c, message)
}
func respondFailure(c *gin.Context, message string) {
response.RespondFailure(c, message)
}
func respondBadRequest(c *gin.Context, message string) {
response.RespondBadRequest(c, message)
}
func respondUnauthorized(c *gin.Context, message string) {
response.RespondUnauthorized(c, message)
}
func decodeJSONBody(body io.Reader, target any) error {
return json.NewDecoder(body).Decode(target)
}
func decodeOptionalJSONBody(body io.Reader, target any) error {
if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) {
return err
}
return nil
}
func parseIDParam(c *gin.Context) (uint, bool) {
return parseIDParamByName(c, "id")
}
func parseIDParamByName(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 {
respondBadRequest(c, "")
return 0, false
}
return uint(id), true
}
func bindJSON(c *gin.Context, target any) bool {
if err := decodeJSONBody(c.Request.Body, target); err != nil {
respondBadRequest(c, "")
return false
}
return true
}
+38 -36
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -16,10 +18,10 @@ import (
func GetTLSCertificates(c *gin.Context) { func GetTLSCertificates(c *gin.Context) {
certificates, err := service.ListTLSCertificates() certificates, err := service.ListTLSCertificates()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificates) response.RespondSuccess(c, certificates)
} }
// GetTLSCertificate godoc // GetTLSCertificate godoc
@@ -32,17 +34,17 @@ func GetTLSCertificates(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id} [get] // @Router /api/tls-certificates/{id} [get]
func GetTLSCertificate(c *gin.Context) { func GetTLSCertificate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
certificate, err := service.GetTLSCertificate(id) certificate, err := service.GetTLSCertificate(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// GetTLSCertificateContent godoc // GetTLSCertificateContent godoc
@@ -55,17 +57,17 @@ func GetTLSCertificate(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/content [get] // @Router /api/tls-certificates/{id}/content [get]
func GetTLSCertificateContent(c *gin.Context) { func GetTLSCertificateContent(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
content, err := service.GetTLSCertificateContent(id) content, err := service.GetTLSCertificateContent(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, content) response.RespondSuccess(c, content)
} }
// CreateTLSCertificate godoc // CreateTLSCertificate godoc
@@ -80,15 +82,15 @@ func GetTLSCertificateContent(c *gin.Context) {
// @Router /api/tls-certificates/ [post] // @Router /api/tls-certificates/ [post]
func CreateTLSCertificate(c *gin.Context) { func CreateTLSCertificate(c *gin.Context) {
var input service.TLSCertificateInput var input service.TLSCertificateInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
certificate, err := service.CreateTLSCertificate(input) certificate, err := service.CreateTLSCertificate(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// UpdateTLSCertificate godoc // UpdateTLSCertificate godoc
@@ -103,22 +105,22 @@ func CreateTLSCertificate(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/update [post] // @Router /api/tls-certificates/{id}/update [post]
func UpdateTLSCertificate(c *gin.Context) { func UpdateTLSCertificate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.TLSCertificateInput var input service.TLSCertificateInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
certificate, err := service.UpdateTLSCertificate(id, input) certificate, err := service.UpdateTLSCertificate(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// ImportTLSCertificateFile godoc // ImportTLSCertificateFile godoc
@@ -139,20 +141,20 @@ func ImportTLSCertificateFile(c *gin.Context) {
remark := c.PostForm("remark") remark := c.PostForm("remark")
certFile, err := c.FormFile("cert_file") certFile, err := c.FormFile("cert_file")
if err != nil { if err != nil {
respondBadRequest(c, "缺少证书文件") response.RespondBadRequest(c, "缺少证书文件")
return return
} }
keyFile, err := c.FormFile("key_file") keyFile, err := c.FormFile("key_file")
if err != nil { if err != nil {
respondBadRequest(c, "缺少私钥文件") response.RespondBadRequest(c, "缺少私钥文件")
return return
} }
certificate, err := service.CreateTLSCertificateFromFiles(name, certFile, keyFile, remark) certificate, err := service.CreateTLSCertificateFromFiles(name, certFile, keyFile, remark)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// DeleteTLSCertificate godoc // DeleteTLSCertificate godoc
@@ -165,15 +167,15 @@ func ImportTLSCertificateFile(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/delete [post] // @Router /api/tls-certificates/{id}/delete [post]
func DeleteTLSCertificate(c *gin.Context) { func DeleteTLSCertificate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteTLSCertificate(id); err != nil { if err := service.DeleteTLSCertificate(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, nil) response.RespondSuccess(c, nil)
} }
// ApplyTLSCertificate godoc // ApplyTLSCertificate godoc
@@ -188,15 +190,15 @@ func DeleteTLSCertificate(c *gin.Context) {
// @Router /api/tls-certificates/apply [post] // @Router /api/tls-certificates/apply [post]
func ApplyTLSCertificate(c *gin.Context) { func ApplyTLSCertificate(c *gin.Context) {
var input service.TLSApplyInput var input service.TLSApplyInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
certificate, err := service.ApplyTLSCertificate(input) certificate, err := service.ApplyTLSCertificate(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// UpdateAcmeCertificate godoc // UpdateAcmeCertificate godoc
@@ -211,21 +213,21 @@ func ApplyTLSCertificate(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/update-acme [post] // @Router /api/tls-certificates/{id}/update-acme [post]
func UpdateAcmeCertificate(c *gin.Context) { func UpdateAcmeCertificate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.TLSApplyInput var input service.TLSApplyInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
certificate, err := service.UpdateAcmeCertificate(id, input) certificate, err := service.UpdateAcmeCertificate(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// ConvertTLSCertificateToAcme godoc // ConvertTLSCertificateToAcme godoc
@@ -240,21 +242,21 @@ func UpdateAcmeCertificate(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/convert-acme [post] // @Router /api/tls-certificates/{id}/convert-acme [post]
func ConvertTLSCertificateToAcme(c *gin.Context) { func ConvertTLSCertificateToAcme(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.TLSApplyInput var input service.TLSApplyInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
certificate, err := service.ConvertTLSCertificateToAcme(id, input) certificate, err := service.ConvertTLSCertificateToAcme(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
// RenewTLSCertificate godoc // RenewTLSCertificate godoc
@@ -267,14 +269,14 @@ func ConvertTLSCertificateToAcme(c *gin.Context) {
// @Failure 400 {object} map[string]interface{} // @Failure 400 {object} map[string]interface{}
// @Router /api/tls-certificates/{id}/renew [post] // @Router /api/tls-certificates/{id}/renew [post]
func RenewTLSCertificate(c *gin.Context) { func RenewTLSCertificate(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
certificate, err := service.RenewTLSCertificate(id) certificate, err := service.RenewTLSCertificate(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, certificate) response.RespondSuccess(c, certificate)
} }
+17 -15
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"time" "time"
@@ -26,10 +28,10 @@ type serverUpgradeRequest struct {
func GetLatestRelease(c *gin.Context) { func GetLatestRelease(c *gin.Context) {
release, err := service.GetLatestServerRelease(c.Request.Context(), c.Query("channel")) release, err := service.GetLatestServerRelease(c.Request.Context(), c.Query("channel"))
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, release) response.RespondSuccess(c, release)
} }
// UpgradeServer godoc // UpgradeServer godoc
@@ -41,18 +43,18 @@ func GetLatestRelease(c *gin.Context) {
func UpgradeServer(c *gin.Context) { func UpgradeServer(c *gin.Context) {
var request serverUpgradeRequest var request serverUpgradeRequest
if c.Request.ContentLength > 0 { if c.Request.ContentLength > 0 {
if err := decodeOptionalJSONBody(c.Request.Body, &request); err != nil { if err := bind.OptionalJSON(c.Request.Body, &request); err != nil {
respondBadRequest(c, "无效的参数") response.RespondBadRequest(c, "无效的参数")
return return
} }
} }
release, err := service.ScheduleServerUpgrade(request.Channel) release, err := service.ScheduleServerUpgrade(request.Channel)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessWithExtras(c, release, gin.H{ response.RespondSuccessWithExtras(c, release, gin.H{
"message": "服务升级任务已启动,下载完成后将自动重启。", "message": "服务升级任务已启动,下载完成后将自动重启。",
}) })
} }
@@ -101,18 +103,18 @@ func StreamServerUpgradeLogs(c *gin.Context) {
// @Success 200 {object} map[string]interface{} // @Success 200 {object} map[string]interface{}
// @Router /api/update/manual-upload [post] // @Router /api/update/manual-upload [post]
func UploadManualServerBinary(c *gin.Context) { func UploadManualServerBinary(c *gin.Context) {
respondFailure(c, "手动升级功能已禁用") response.RespondFailure(c, "手动升级功能已禁用")
return return
// //
//fileHeader, err := c.FormFile("binary") //fileHeader, err := c.FormFile("binary")
//if err != nil { //if err != nil {
// respondFailure(c, "请先选择要上传的服务端二进制文件。") // response.RespondFailure(c, "请先选择要上传的服务端二进制文件。")
// return // return
//} //}
// //
//file, err := fileHeader.Open() //file, err := fileHeader.Open()
//if err != nil { //if err != nil {
// respondFailure(c, "读取上传文件失败。") // response.RespondFailure(c, "读取上传文件失败。")
// return // return
//} //}
//defer func() { //defer func() {
@@ -121,7 +123,7 @@ func UploadManualServerBinary(c *gin.Context) {
// //
//info, err := service.UploadManualServerBinary(c.Request.Context(), fileHeader.Filename, file) //info, err := service.UploadManualServerBinary(c.Request.Context(), fileHeader.Filename, file)
//if err != nil { //if err != nil {
// respondFailure(c, err.Error()) // response.RespondFailure(c, err.Error())
// return // return
//} //}
// //
@@ -130,7 +132,7 @@ func UploadManualServerBinary(c *gin.Context) {
// message = "已完成上传并检查升级包版本。" // message = "已完成上传并检查升级包版本。"
//} //}
// //
//respondSuccessWithExtras(c, info, gin.H{ //response.RespondSuccessWithExtras(c, info, gin.H{
// "message": message, // "message": message,
//}) //})
} }
@@ -143,21 +145,21 @@ func UploadManualServerBinary(c *gin.Context) {
// @Success 200 {object} map[string]interface{} // @Success 200 {object} map[string]interface{}
// @Router /api/update/manual-upgrade [post] // @Router /api/update/manual-upgrade [post]
func ConfirmManualServerUpgrade(c *gin.Context) { func ConfirmManualServerUpgrade(c *gin.Context) {
respondFailure(c, "手动升级功能已禁用") response.RespondFailure(c, "手动升级功能已禁用")
return return
// //
//var request confirmManualUpgradeRequest //var request confirmManualUpgradeRequest
//if !bindJSON(c, &request) { //if !bind.JSON(c, &request) {
// return // return
//} //}
// //
//info, err := service.ConfirmManualServerUpgrade(request.UploadToken) //info, err := service.ConfirmManualServerUpgrade(request.UploadToken)
//if err != nil { //if err != nil {
// respondFailure(c, err.Error()) // response.RespondFailure(c, err.Error())
// return // return
//} //}
// //
//respondSuccessWithExtras(c, info, gin.H{ //response.RespondSuccessWithExtras(c, info, gin.H{
// "message": "服务升级任务已启动,确认无误后将自动重启。", // "message": "服务升级任务已启动,确认无误后将自动重启。",
//}) //})
} }
+3 -2
View File
@@ -1,6 +1,7 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/service" "openflare/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -17,8 +18,8 @@ import (
func SyncUptimeKuma(c *gin.Context) { func SyncUptimeKuma(c *gin.Context) {
err := service.SyncToUptimeKuma() err := service.SyncToUptimeKuma()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "同步成功") response.RespondSuccessMessage(c, "同步成功")
} }
+65 -63
View File
@@ -2,6 +2,8 @@ package controller
import ( import (
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/controller/bind"
"openflare/middleware" "openflare/middleware"
"openflare/model" "openflare/model"
"openflare/utils/security" "openflare/utils/security"
@@ -18,17 +20,17 @@ type LoginRequest struct {
func Login(c *gin.Context) { func Login(c *gin.Context) {
if !common.PasswordLoginEnabled { if !common.PasswordLoginEnabled {
respondFailure(c, "管理员关闭了密码登录") response.RespondFailure(c, "管理员关闭了密码登录")
return return
} }
var loginRequest LoginRequest var loginRequest LoginRequest
if !bindJSON(c, &loginRequest) { if !bind.JSON(c, &loginRequest) {
return return
} }
username := loginRequest.Username username := loginRequest.Username
password := loginRequest.Password password := loginRequest.Password
if username == "" || password == "" { if username == "" || password == "" {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
user := model.User{ user := model.User{
@@ -37,7 +39,7 @@ func Login(c *gin.Context) {
} }
err := user.ValidateAndFill() err := user.ValidateAndFill()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
setupLogin(&user, c) setupLogin(&user, c)
@@ -68,10 +70,10 @@ func setLoginToken(user *model.User) (*model.User, error) {
func setupLogin(user *model.User, c *gin.Context) { func setupLogin(user *model.User, c *gin.Context) {
cleanUser, err := setLoginToken(user) cleanUser, err := setLoginToken(user)
if err != nil { if err != nil {
respondFailure(c, "无法保存会话信息,请重试") response.RespondFailure(c, "无法保存会话信息,请重试")
return return
} }
respondSuccess(c, *cleanUser) response.RespondSuccess(c, *cleanUser)
} }
func Logout(c *gin.Context) { func Logout(c *gin.Context) {
@@ -80,12 +82,12 @@ func Logout(c *gin.Context) {
user := model.ValidateUserToken(token) user := model.ValidateUserToken(token)
if user != nil && user.Id != 0 { if user != nil && user.Id != 0 {
if err := model.DB.Model(user).Update("token", "").Error; err != nil { if err := model.DB.Model(user).Update("token", "").Error; err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
} }
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func currentUserFromOpenFlareToken(c *gin.Context) *model.User { func currentUserFromOpenFlareToken(c *gin.Context) *model.User {
@@ -97,7 +99,7 @@ func currentUserFromOpenFlareToken(c *gin.Context) *model.User {
} }
func Register(c *gin.Context) { func Register(c *gin.Context) {
respondFailure(c, "非法请求") response.RespondFailure(c, "非法请求")
} }
func GetAllUsers(c *gin.Context) { func GetAllUsers(c *gin.Context) {
@@ -107,99 +109,99 @@ func GetAllUsers(c *gin.Context) {
} }
users, err := model.GetAllUsers(p*common.ItemsPerPage, common.ItemsPerPage) users, err := model.GetAllUsers(p*common.ItemsPerPage, common.ItemsPerPage)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, users) response.RespondSuccess(c, users)
} }
func SearchUsers(c *gin.Context) { func SearchUsers(c *gin.Context) {
keyword := c.Query("keyword") keyword := c.Query("keyword")
users, err := model.SearchUsers(keyword) users, err := model.SearchUsers(keyword)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, users) response.RespondSuccess(c, users)
} }
func GetUser(c *gin.Context) { func GetUser(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
user, err := model.GetUserById(int(id), false) user, err := model.GetUserById(int(id), false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
if myRole <= user.Role { if myRole <= user.Role {
respondFailure(c, "无权获取同级或更高等级用户的信息") response.RespondFailure(c, "无权获取同级或更高等级用户的信息")
return return
} }
respondSuccess(c, user) response.RespondSuccess(c, user)
} }
func GenerateToken(c *gin.Context) { func GenerateToken(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, true) user, err := model.GetUserById(id, true)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
// Generate a fresh JWT for the user // Generate a fresh JWT for the user
tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user) tokenString, _, err := middleware.JWTMiddleware.TokenGenerator(user)
if err != nil { if err != nil {
respondFailure(c, "生成 Token 失败: "+err.Error()) response.RespondFailure(c, "生成 Token 失败: "+err.Error())
return return
} }
user.Token = tokenString user.Token = tokenString
if err := user.Update(false); err != nil { if err := user.Update(false); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, user.Token) response.RespondSuccess(c, user.Token)
} }
func GetSelf(c *gin.Context) { func GetSelf(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
user, err := model.GetUserById(id, false) user, err := model.GetUserById(id, false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, user) response.RespondSuccess(c, user)
} }
func UpdateUser(c *gin.Context) { func UpdateUser(c *gin.Context) {
var updatedUser model.User var updatedUser model.User
if !bindJSON(c, &updatedUser) { if !bind.JSON(c, &updatedUser) {
return return
} }
if updatedUser.Id == 0 { if updatedUser.Id == 0 {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
if updatedUser.Password == "" { if updatedUser.Password == "" {
updatedUser.Password = "$I_LOVE_U" // make Validator happy :) updatedUser.Password = "$I_LOVE_U" // make Validator happy :)
} }
if err := validation.Validate.Struct(&updatedUser); err != nil { if err := validation.Validate.Struct(&updatedUser); err != nil {
respondFailure(c, "输入不合法 "+err.Error()) response.RespondFailure(c, "输入不合法 "+err.Error())
return return
} }
originUser, err := model.GetUserById(updatedUser.Id, false) originUser, err := model.GetUserById(updatedUser.Id, false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
if myRole <= originUser.Role { if myRole <= originUser.Role {
respondFailure(c, "无权更新同权限等级或更高权限等级的用户信息") response.RespondFailure(c, "无权更新同权限等级或更高权限等级的用户信息")
return return
} }
if myRole <= updatedUser.Role { if myRole <= updatedUser.Role {
respondFailure(c, "无权将其他用户权限等级提升到大于等于自己的权限等级") response.RespondFailure(c, "无权将其他用户权限等级提升到大于等于自己的权限等级")
return return
} }
if updatedUser.Password == "$I_LOVE_U" { if updatedUser.Password == "$I_LOVE_U" {
@@ -207,22 +209,22 @@ func UpdateUser(c *gin.Context) {
} }
updatePassword := updatedUser.Password != "" updatePassword := updatedUser.Password != ""
if err := updatedUser.Update(updatePassword); err != nil { if err := updatedUser.Update(updatePassword); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func UpdateSelf(c *gin.Context) { func UpdateSelf(c *gin.Context) {
var user model.User var user model.User
if !bindJSON(c, &user) { if !bind.JSON(c, &user) {
return return
} }
if user.Password == "" { if user.Password == "" {
user.Password = "$I_LOVE_U" // make Validator happy :) user.Password = "$I_LOVE_U" // make Validator happy :)
} }
if err := validation.Validate.Struct(&user); err != nil { if err := validation.Validate.Struct(&user); err != nil {
respondFailure(c, "输入不合法 "+err.Error()) response.RespondFailure(c, "输入不合法 "+err.Error())
return return
} }
@@ -238,53 +240,53 @@ func UpdateSelf(c *gin.Context) {
} }
updatePassword := user.Password != "" updatePassword := user.Password != ""
if err := cleanUser.Update(updatePassword); err != nil { if err := cleanUser.Update(updatePassword); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func DeleteUser(c *gin.Context) { func DeleteUser(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
originUser, err := model.GetUserById(int(id), false) originUser, err := model.GetUserById(int(id), false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
if myRole <= originUser.Role { if myRole <= originUser.Role {
respondFailure(c, "无权删除同权限等级或更高权限等级的用户") response.RespondFailure(c, "无权删除同权限等级或更高权限等级的用户")
return return
} }
err = model.DeleteUserById(int(id)) err = model.DeleteUserById(int(id))
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func DeleteSelf(c *gin.Context) { func DeleteSelf(c *gin.Context) {
id := c.GetInt("id") id := c.GetInt("id")
err := model.DeleteUserById(id) err := model.DeleteUserById(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func CreateUser(c *gin.Context) { func CreateUser(c *gin.Context) {
var user model.User var user model.User
if !bindJSON(c, &user) { if !bind.JSON(c, &user) {
return return
} }
if user.Username == "" || user.Password == "" { if user.Username == "" || user.Password == "" {
respondFailure(c, "无效的参数") response.RespondFailure(c, "无效的参数")
return return
} }
if user.DisplayName == "" { if user.DisplayName == "" {
@@ -292,7 +294,7 @@ func CreateUser(c *gin.Context) {
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
if user.Role >= myRole { if user.Role >= myRole {
respondFailure(c, "无法创建权限大于等于自己的用户") response.RespondFailure(c, "无法创建权限大于等于自己的用户")
return return
} }
// Even for admin users, we cannot fully trust them! // Even for admin users, we cannot fully trust them!
@@ -302,11 +304,11 @@ func CreateUser(c *gin.Context) {
DisplayName: user.DisplayName, DisplayName: user.DisplayName,
} }
if err := cleanUser.Insert(); err != nil { if err := cleanUser.Insert(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
type ManageRequest struct { type ManageRequest struct {
@@ -317,7 +319,7 @@ type ManageRequest struct {
// ManageUser Only admin user can do this // ManageUser Only admin user can do this
func ManageUser(c *gin.Context) { func ManageUser(c *gin.Context) {
var req ManageRequest var req ManageRequest
if !bindJSON(c, &req) { if !bind.JSON(c, &req) {
return return
} }
user := model.User{ user := model.User{
@@ -326,70 +328,70 @@ func ManageUser(c *gin.Context) {
// Fill attributes // Fill attributes
model.DB.Where(&user).First(&user) model.DB.Where(&user).First(&user)
if user.Id == 0 { if user.Id == 0 {
respondFailure(c, "用户不存在") response.RespondFailure(c, "用户不存在")
return return
} }
myRole := c.GetInt("role") myRole := c.GetInt("role")
if myRole <= user.Role && myRole != common.RoleRootUser { if myRole <= user.Role && myRole != common.RoleRootUser {
respondFailure(c, "无权更新同权限等级或更高权限等级的用户信息") response.RespondFailure(c, "无权更新同权限等级或更高权限等级的用户信息")
return return
} }
switch req.Action { switch req.Action {
case "disable": case "disable":
user.Status = common.UserStatusDisabled user.Status = common.UserStatusDisabled
if user.Role == common.RoleRootUser { if user.Role == common.RoleRootUser {
respondFailure(c, "无法禁用超级管理员用户") response.RespondFailure(c, "无法禁用超级管理员用户")
return return
} }
case "enable": case "enable":
user.Status = common.UserStatusEnabled user.Status = common.UserStatusEnabled
case "delete": case "delete":
if user.Role == common.RoleRootUser { if user.Role == common.RoleRootUser {
respondFailure(c, "无法删除超级管理员用户") response.RespondFailure(c, "无法删除超级管理员用户")
return return
} }
if err := user.Delete(); err != nil { if err := user.Delete(); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
case "promote": case "promote":
if myRole != common.RoleRootUser { if myRole != common.RoleRootUser {
respondFailure(c, "普通管理员用户无法提升其他用户为管理员") response.RespondFailure(c, "普通管理员用户无法提升其他用户为管理员")
return return
} }
if user.Role >= common.RoleAdminUser { if user.Role >= common.RoleAdminUser {
respondFailure(c, "该用户已经是管理员") response.RespondFailure(c, "该用户已经是管理员")
return return
} }
user.Role = common.RoleAdminUser user.Role = common.RoleAdminUser
case "demote": case "demote":
if user.Role == common.RoleRootUser { if user.Role == common.RoleRootUser {
respondFailure(c, "无法降级超级管理员用户") response.RespondFailure(c, "无法降级超级管理员用户")
return return
} }
if user.Role == common.RoleCommonUser { if user.Role == common.RoleCommonUser {
respondFailure(c, "该用户已经是普通用户") response.RespondFailure(c, "该用户已经是普通用户")
return return
} }
user.Role = common.RoleCommonUser user.Role = common.RoleCommonUser
} }
if err := user.Update(false); err != nil { if err := user.Update(false); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
clearUser := model.User{ clearUser := model.User{
Role: user.Role, Role: user.Role,
Status: user.Status, Status: user.Status,
} }
respondSuccess(c, clearUser) response.RespondSuccess(c, clearUser)
} }
func EmailBind(c *gin.Context) { func EmailBind(c *gin.Context) {
email := c.Query("email") email := c.Query("email")
code := c.Query("code") code := c.Query("code")
if !security.VerifyCodeWithKey(email, code, security.EmailVerificationPurpose) { if !security.VerifyCodeWithKey(email, code, security.EmailVerificationPurpose) {
respondFailure(c, "验证码错误或已过期") response.RespondFailure(c, "验证码错误或已过期")
return return
} }
id := c.GetInt("id") id := c.GetInt("id")
@@ -398,15 +400,15 @@ func EmailBind(c *gin.Context) {
} }
err := user.FillUserById() err := user.FillUserById()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user.Email = email user.Email = email
// no need to check if this email already taken, because we have used verification code to check it // no need to check if this email already taken, because we have used verification code to check it
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
+48 -46
View File
@@ -1,6 +1,8 @@
package controller package controller
import ( import (
"openflare/common/response"
"openflare/controller/bind"
"openflare/service" "openflare/service"
"strconv" "strconv"
@@ -14,82 +16,82 @@ type wafIDsRequest struct {
func ListWAFRuleGroups(c *gin.Context) { func ListWAFRuleGroups(c *gin.Context) {
groups, err := service.ListWAFRuleGroups() groups, err := service.ListWAFRuleGroups()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, groups) response.RespondSuccess(c, groups)
} }
func GetWAFRuleGroup(c *gin.Context) { func GetWAFRuleGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
group, err := service.GetWAFRuleGroup(id) group, err := service.GetWAFRuleGroup(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func CreateWAFRuleGroup(c *gin.Context) { func CreateWAFRuleGroup(c *gin.Context) {
var input service.WAFRuleGroupInput var input service.WAFRuleGroupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
group, err := service.CreateWAFRuleGroup(input) group, err := service.CreateWAFRuleGroup(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func UpdateWAFRuleGroup(c *gin.Context) { func UpdateWAFRuleGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.WAFRuleGroupInput var input service.WAFRuleGroupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
group, err := service.UpdateWAFRuleGroup(id, input) group, err := service.UpdateWAFRuleGroup(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func DeleteWAFRuleGroup(c *gin.Context) { func DeleteWAFRuleGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteWAFRuleGroup(id); err != nil { if err := service.DeleteWAFRuleGroup(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func ReplaceWAFRuleGroupSites(c *gin.Context) { func ReplaceWAFRuleGroupSites(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var request wafIDsRequest var request wafIDsRequest
if !bindJSON(c, &request) { if !bind.JSON(c, &request) {
return return
} }
group, err := service.ReplaceWAFRuleGroupSites(id, request.IDs) group, err := service.ReplaceWAFRuleGroupSites(id, request.IDs)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func GetWAFSiteRuleGroups(c *gin.Context) { func GetWAFSiteRuleGroups(c *gin.Context) {
@@ -99,10 +101,10 @@ func GetWAFSiteRuleGroups(c *gin.Context) {
} }
view, err := service.GetWAFSiteRuleGroups(routeID) view, err := service.GetWAFSiteRuleGroups(routeID)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, view) response.RespondSuccess(c, view)
} }
func ReplaceWAFSiteRuleGroups(c *gin.Context) { func ReplaceWAFSiteRuleGroups(c *gin.Context) {
@@ -111,111 +113,111 @@ func ReplaceWAFSiteRuleGroups(c *gin.Context) {
return return
} }
var request wafIDsRequest var request wafIDsRequest
if !bindJSON(c, &request) { if !bind.JSON(c, &request) {
return return
} }
view, err := service.ReplaceWAFSiteRuleGroups(routeID, request.IDs) view, err := service.ReplaceWAFSiteRuleGroups(routeID, request.IDs)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, view) response.RespondSuccess(c, view)
} }
func ListWAFIPGroups(c *gin.Context) { func ListWAFIPGroups(c *gin.Context) {
groups, err := service.ListWAFIPGroups() groups, err := service.ListWAFIPGroups()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, groups) response.RespondSuccess(c, groups)
} }
func GetWAFIPGroup(c *gin.Context) { func GetWAFIPGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
group, err := service.GetWAFIPGroup(id) group, err := service.GetWAFIPGroup(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func CreateWAFIPGroup(c *gin.Context) { func CreateWAFIPGroup(c *gin.Context) {
var input service.WAFIPGroupInput var input service.WAFIPGroupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
group, err := service.CreateWAFIPGroup(input) group, err := service.CreateWAFIPGroup(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func UpdateWAFIPGroup(c *gin.Context) { func UpdateWAFIPGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
var input service.WAFIPGroupInput var input service.WAFIPGroupInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
group, err := service.UpdateWAFIPGroup(id, input) group, err := service.UpdateWAFIPGroup(id, input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, group) response.RespondSuccess(c, group)
} }
func DeleteWAFIPGroup(c *gin.Context) { func DeleteWAFIPGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
if err := service.DeleteWAFIPGroup(id); err != nil { if err := service.DeleteWAFIPGroup(id); err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
} }
func SyncWAFIPGroup(c *gin.Context) { func SyncWAFIPGroup(c *gin.Context) {
id, ok := parseIDParam(c) id, ok := bind.IDParam(c)
if !ok { if !ok {
return return
} }
result, err := service.SyncWAFIPGroup(id) result, err := service.SyncWAFIPGroup(id)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
func TestWAFIPGroupAutoConfig(c *gin.Context) { func TestWAFIPGroupAutoConfig(c *gin.Context) {
var input service.WAFIPGroupAutoTestInput var input service.WAFIPGroupAutoTestInput
if !bindJSON(c, &input) { if !bind.JSON(c, &input) {
return return
} }
result, err := service.TestWAFIPGroupAutoConfig(input) result, err := service.TestWAFIPGroupAutoConfig(input)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccess(c, result) response.RespondSuccess(c, result)
} }
func parseUintPathParam(c *gin.Context, name string) (uint, bool) { func parseUintPathParam(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64) id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 { if err != nil || id == 0 {
respondBadRequest(c, "invalid id") response.RespondBadRequest(c, "invalid id")
return 0, false return 0, false
} }
return uint(id), true return uint(id), true
+12 -11
View File
@@ -8,6 +8,7 @@ import (
"log/slog" "log/slog"
"net/http" "net/http"
"openflare/common" "openflare/common"
"openflare/common/response"
"openflare/model" "openflare/model"
"time" "time"
@@ -58,13 +59,13 @@ func getWeChatIdByCode(code string) (string, error) {
func WeChatAuth(c *gin.Context) { func WeChatAuth(c *gin.Context) {
if !common.WeChatAuthEnabled { if !common.WeChatAuthEnabled {
respondFailure(c, "管理员未开启通过微信登录以及注册") response.RespondFailure(c, "管理员未开启通过微信登录以及注册")
return return
} }
code := c.Query("code") code := c.Query("code")
wechatId, err := getWeChatIdByCode(code) wechatId, err := getWeChatIdByCode(code)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user := model.User{ user := model.User{
@@ -73,16 +74,16 @@ func WeChatAuth(c *gin.Context) {
if model.IsWeChatIdAlreadyTaken(wechatId) { if model.IsWeChatIdAlreadyTaken(wechatId) {
err := user.FillUserByWeChatId() err := user.FillUserByWeChatId()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
} else { } else {
respondFailure(c, "管理员关闭了新用户注册") response.RespondFailure(c, "管理员关闭了新用户注册")
return return
} }
if user.Status != common.UserStatusEnabled { if user.Status != common.UserStatusEnabled {
respondFailure(c, "用户已被封禁") response.RespondFailure(c, "用户已被封禁")
return return
} }
setupLogin(&user, c) setupLogin(&user, c)
@@ -90,17 +91,17 @@ func WeChatAuth(c *gin.Context) {
func WeChatBind(c *gin.Context) { func WeChatBind(c *gin.Context) {
if !common.WeChatAuthEnabled { if !common.WeChatAuthEnabled {
respondFailure(c, "管理员未开启通过微信登录以及注册") response.RespondFailure(c, "管理员未开启通过微信登录以及注册")
return return
} }
code := c.Query("code") code := c.Query("code")
wechatId, err := getWeChatIdByCode(code) wechatId, err := getWeChatIdByCode(code)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
if model.IsWeChatIdAlreadyTaken(wechatId) { if model.IsWeChatIdAlreadyTaken(wechatId) {
respondFailure(c, "该微信账号已被绑定") response.RespondFailure(c, "该微信账号已被绑定")
return return
} }
id := c.GetInt("id") id := c.GetInt("id")
@@ -109,15 +110,15 @@ func WeChatBind(c *gin.Context) {
} }
err = user.FillUserById() err = user.FillUserById()
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
user.WeChatId = wechatId user.WeChatId = wechatId
err = user.Update(false) err = user.Update(false)
if err != nil { if err != nil {
respondFailure(c, err.Error()) response.RespondFailure(c, err.Error())
return return
} }
respondSuccessMessage(c, "") response.RespondSuccessMessage(c, "")
return return
} }