Files
ai_site/platform/internal/handler/agent_sync.go
whm bb59bb5efd fix: allow agent JWT to use attached online_db_id (Z14c)
Agent tokens use agent_id as user_id while Binding is under the human member; authorize schema/push/ensure against the agent's mounted OnlineDBID.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-05 18:22:57 +08:00

503 lines
17 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package handler
import (
"encoding/json"
"log"
"net/http"
"strings"
"time"
"aijianzhan/platform/internal/audit"
"aijianzhan/platform/internal/authx"
"aijianzhan/platform/internal/dbsync"
"aijianzhan/platform/internal/svc"
"github.com/zeromicro/go-zero/rest/httpx"
"github.com/zeromicro/go-zero/rest/pathvar"
)
// agent 同步:管理员/智能体(「数据同步」)租户级;普通登录用户按本人 Binding 范围。
func agentSyncWhitelistHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
if !authorizeUserSyncChannel(svcCtx, w, r, channelID) {
return
}
tables := uniqueStringSlice(ch.Local.Tables, ch.Remote.Tables)
hint := "通道 tables 仅历史兼容;用户自助以 Binding+JWT 为权限,整库表均可 push可不改白名单"
if len(tables) == 0 {
hint = "通道表白名单为空;以 Binding 为准,接受任意表 push可自动建表"
}
httpx.OkJson(w, map[string]any{
"channel_id": ch.ID,
"name": ch.Name,
"enabled": ch.Enabled,
"direction": ch.Direction,
"conflict_policy": ch.ConflictPolicy,
"tables": tables,
"pk_columns": ch.PKColumns,
"hint": hint,
})
}
}
func agentSyncPushHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var item dbsync.PushItem
if err := json.NewDecoder(r.Body).Decode(&item); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
if !authorizeUserSyncPush(svcCtx, w, r, channelID, item) {
return
}
res, err := dbsync.PushToRemote(r.Context(), ch, svcCtx.DBSync.Store(), item)
dur := time.Since(started)
if err != nil {
logSyncReq(r, "push", channelID, item.Table, item.RowPK, "err", err.Error(), dur, reqID)
auditSync(svcCtx, r, "dbsync.push.error", map[string]any{
"channel_id": channelID, "table": item.Table, "row_pk": item.RowPK,
"online_db_id": item.OnlineDBID, "error": err.Error(), "ms": dur.Milliseconds(), "req_id": reqID,
})
if dbsync.IsRetryable(err) {
authx.WriteRetryableError(w, http.StatusServiceUnavailable, err.Error())
return
}
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
outcome := "applied"
if res != nil && res.Skipped {
outcome = "skipped"
}
logSyncReq(r, "push", channelID, item.Table, item.RowPK, outcome, "", dur, reqID)
auditSync(svcCtx, r, "dbsync.push", map[string]any{
"channel_id": channelID, "table": item.Table, "row_pk": item.RowPK,
"online_db_id": item.OnlineDBID, "outcome": outcome, "ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "result": res})
}
}
func agentSyncPushBatchHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var body struct {
Items []dbsync.PushItem `json:"items"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
if len(body.Items) == 0 {
authx.WriteError(w, http.StatusBadRequest, "items required")
return
}
if len(body.Items) > 100 {
authx.WriteError(w, http.StatusBadRequest, "items limit 100")
return
}
for _, item := range body.Items {
if !authorizeUserSyncPush(svcCtx, w, r, channelID, item) {
return
}
}
results, err := dbsync.PushBatchToRemote(r.Context(), ch, svcCtx.DBSync.Store(), body.Items)
dur := time.Since(started)
if err != nil {
logSyncReq(r, "push_batch", channelID, "", "", "err", err.Error(), dur, reqID)
auditSync(svcCtx, r, "dbsync.push_batch.error", map[string]any{
"channel_id": channelID, "n": len(body.Items), "error": err.Error(),
"ms": dur.Milliseconds(), "req_id": reqID,
})
payload := map[string]any{
"success": false,
"error": err.Error(),
"results": results,
}
if dbsync.IsRetryable(err) {
payload["retryable"] = true
httpx.WriteJson(w, http.StatusServiceUnavailable, payload)
return
}
httpx.OkJson(w, payload)
return
}
logSyncReq(r, "push_batch", channelID, "", "", "ok", "", dur, reqID)
auditSync(svcCtx, r, "dbsync.push_batch", map[string]any{
"channel_id": channelID, "n": len(body.Items), "ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "results": results})
}
}
func logSyncReq(r *http.Request, op, channelID, table, rowPK, outcome, errMsg string, dur time.Duration, reqID string) {
uid := authx.UserID(r.Context())
if errMsg != "" {
log.Printf("dbsync %s channel=%s table=%s pk=%s user=%d outcome=%s err=%s dur=%s req=%s",
op, channelID, table, rowPK, uid, outcome, errMsg, dur, reqID)
return
}
log.Printf("dbsync %s channel=%s table=%s pk=%s user=%d outcome=%s dur=%s req=%s",
op, channelID, table, rowPK, uid, outcome, dur, reqID)
}
func auditSync(svcCtx *svc.ServiceContext, r *http.Request, action string, detail map[string]any) {
if svcCtx == nil || svcCtx.Audit == nil {
return
}
_ = svcCtx.Audit.Log(r.Context(), syncTenantID(r), authx.UserID(r.Context()), action, audit.DetailJSON(detail))
}
// authorizeUserSyncChannel 无「数据同步」时,须本人 Binding 覆盖该通道。
func authorizeUserSyncChannel(svcCtx *svc.ServiceContext, w http.ResponseWriter, r *http.Request, channelID string) bool {
if authx.SyncTenantWide(r.Context()) {
return true
}
uid := authx.UserID(r.Context())
tid := syncTenantID(r)
if !svcCtx.DBSync.Store().UserCanAccessChannel(tid, uid, channelID) {
authx.WriteError(w, http.StatusForbidden, "无权访问该同步通道:请先登记本人 Bindinglocal_database_id ↔ online_db_id + channel_id")
return false
}
return true
}
// authorizeUserSyncPush 无「数据同步」时,须本人 Binding 且 online_db_id 匹配。
func authorizeUserSyncPush(svcCtx *svc.ServiceContext, w http.ResponseWriter, r *http.Request, channelID string, item dbsync.PushItem) bool {
return authorizeUserOnlineDB(svcCtx, w, r, channelID, item.OnlineDBID, "用户自助 push 须带 online_db_id且须为本人 Binding")
}
// authorizeUserOnlineDB带 online_db_id 时须本人 Binding或智能体挂载落点匹配Z14c
// 仅「数据同步」且未带 online_db_id 时保持租户级管理路径。
func authorizeUserOnlineDB(svcCtx *svc.ServiceContext, w http.ResponseWriter, r *http.Request, channelID, onlineDBID, emptyMsg string) bool {
online := strings.TrimSpace(onlineDBID)
wide := authx.SyncTenantWide(r.Context())
if online == "" {
if wide {
return true
}
if emptyMsg == "" {
emptyMsg = "用户自助须带 online_db_id且须为本人 Binding"
}
authx.WriteError(w, http.StatusForbidden, emptyMsg)
return false
}
uid := authx.UserID(r.Context())
tid := syncTenantID(r)
if svcCtx.DBSync.Store().UserOwnsOnlineDB(tid, uid, channelID, online) {
return true
}
// Z14c智能体 JWT 的 user_id=agent_idBinding 记在人类成员上;放行 agents/me / ticket 挂载的 online_db_id
if agentOwnsOnlineDB(svcCtx, r, tid, channelID, online) {
return true
}
authx.WriteError(w, http.StatusForbidden, "无权访问该 online_db_id非本人 Binding / 非智能体挂载落点 / 非公司共享库)")
return false
}
// agentOwnsOnlineDB 当前智能体账号是否已挂载该 online_db_id及可选 channel
func agentOwnsOnlineDB(svcCtx *svc.ServiceContext, r *http.Request, tenantID int64, channelID, onlineDBID string) bool {
if svcCtx == nil || svcCtx.Agents == nil {
return false
}
aid := authx.AgentID(r.Context())
if aid <= 0 {
// 兼容旧票Role=智能体 时 UserID 即 AgentID
if authx.Role(r.Context()) == authx.Role智能体 {
aid = authx.UserID(r.Context())
}
}
if aid <= 0 || tenantID <= 0 {
return false
}
acc, err := svcCtx.Agents.Get(r.Context(), tenantID, aid)
if err != nil || acc == nil {
return false
}
if strings.TrimSpace(acc.OnlineDBID) != strings.TrimSpace(onlineDBID) {
return false
}
ch := strings.TrimSpace(channelID)
if ch != "" && strings.TrimSpace(acc.ChannelID) != "" && strings.TrimSpace(acc.ChannelID) != ch {
return false
}
return true
}
func agentSyncPullHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return agentSyncPullWithDefaultMode(svcCtx, "")
}
func agentSyncBootstrapHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return agentSyncPullWithDefaultMode(svcCtx, dbsync.PullModeBootstrap)
}
func agentSyncPullWithDefaultMode(svcCtx *svc.ServiceContext, forceMode string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var req dbsync.PullRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
if forceMode != "" {
req.Mode = forceMode
}
if !authorizeUserOnlineDB(svcCtx, w, r, channelID, req.OnlineDBID, "用户自助 pull/bootstrap 须带 online_db_id且须为本人 Binding") {
return
}
if !authorizeUserSyncChannel(svcCtx, w, r, channelID) {
return
}
res, err := dbsync.PullFromRemote(r.Context(), ch, svcCtx.DBSync.Store(), req)
dur := time.Since(started)
mode := req.Mode
if mode == "" {
mode = dbsync.PullModeBootstrap
}
if err != nil {
logSyncReq(r, "pull/"+mode, channelID, req.Table, "", "err", err.Error(), dur, reqID)
auditSync(svcCtx, r, "dbsync.pull.error", map[string]any{
"channel_id": channelID, "mode": mode, "table": req.Table,
"error": err.Error(), "ms": dur.Milliseconds(), "req_id": reqID,
})
if dbsync.IsRetryable(err) {
authx.WriteRetryableError(w, http.StatusServiceUnavailable, err.Error())
return
}
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
n := 0
if res != nil {
n = len(res.Items) + len(res.PKs)
}
logSyncReq(r, "pull/"+mode, channelID, req.Table, "", "ok", "", dur, reqID)
auditSync(svcCtx, r, "dbsync.pull", map[string]any{
"channel_id": channelID, "mode": mode, "table": req.Table,
"n": n, "ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "result": res})
}
}
func agentSyncSchemaEnsureHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var req dbsync.SchemaEnsureRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
online := strings.TrimSpace(req.OnlineDBID)
if online == "" && len(req.Tables) > 0 {
online = strings.TrimSpace(req.Tables[0].OnlineDBID)
req.OnlineDBID = online
}
if !authorizeUserOnlineDB(svcCtx, w, r, channelID, online, "用户自助 schema/ensure 须带 online_db_id且须为本人 Binding") {
return
}
if !authorizeUserSyncChannel(svcCtx, w, r, channelID) {
return
}
res, err := dbsync.EnsureSchemasOnRemote(r.Context(), ch, req)
dur := time.Since(started)
if err != nil {
logSyncReq(r, "schema/ensure", channelID, "", "", "err", err.Error(), dur, reqID)
auditSync(svcCtx, r, "dbsync.schema_ensure.error", map[string]any{
"channel_id": channelID, "n": len(req.Tables), "error": err.Error(),
"ms": dur.Milliseconds(), "req_id": reqID,
})
if dbsync.IsRetryable(err) {
authx.WriteRetryableError(w, http.StatusServiceUnavailable, err.Error())
return
}
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
n := 0
created := 0
if res != nil {
n = len(res.Results)
for _, it := range res.Results {
if it.Created {
created++
}
}
}
logSyncReq(r, "schema/ensure", channelID, "", "", "ok", "", dur, reqID)
auditSync(svcCtx, r, "dbsync.schema_ensure", map[string]any{
"channel_id": channelID, "n": n, "created": created,
"ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "result": res})
}
}
func agentSyncSchemaDescribeHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var req dbsync.SchemaDescribeRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
if !authorizeUserOnlineDB(svcCtx, w, r, channelID, req.OnlineDBID, "用户自助 schema 须带 online_db_id且须为本人 Binding") {
return
}
if !authorizeUserSyncChannel(svcCtx, w, r, channelID) {
return
}
res, err := dbsync.DescribeSchemasFromRemote(r.Context(), ch, req)
dur := time.Since(started)
if err != nil {
logSyncReq(r, "schema", channelID, "", "", "err", err.Error(), dur, reqID)
auditSync(svcCtx, r, "dbsync.schema.error", map[string]any{
"channel_id": channelID, "error": err.Error(),
"ms": dur.Milliseconds(), "req_id": reqID,
})
if dbsync.IsRetryable(err) {
authx.WriteRetryableError(w, http.StatusServiceUnavailable, err.Error())
return
}
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
n := 0
if res != nil {
n = len(res.Tables)
}
logSyncReq(r, "schema", channelID, "", "", "ok", "", dur, reqID)
auditSync(svcCtx, r, "dbsync.schema", map[string]any{
"channel_id": channelID, "n": n, "ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "result": res})
}
}
func agentSyncSchemaFingerprintHandler(svcCtx *svc.ServiceContext) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
started := time.Now()
reqID := r.Header.Get("X-Request-Id")
if !requireDBSync(svcCtx, w) {
return
}
channelID := pathvar.Vars(r)["id"]
ch, err := svcCtx.DBSync.Store().GetChannelForTenant(channelID, syncTenantID(r))
if err != nil {
authx.WriteError(w, http.StatusNotFound, err.Error())
return
}
var req dbsync.FingerprintRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
if !authorizeUserOnlineDB(svcCtx, w, r, channelID, req.OnlineDBID, "用户自助 schema/fingerprint 须带 online_db_id且须为本人 Binding") {
return
}
if !authorizeUserSyncChannel(svcCtx, w, r, channelID) {
return
}
res, err := dbsync.FingerprintRemoteTables(r.Context(), ch, req)
dur := time.Since(started)
if err != nil {
logSyncReq(r, "schema/fingerprint", channelID, "", "", "err", err.Error(), dur, reqID)
if dbsync.IsRetryable(err) {
authx.WriteRetryableError(w, http.StatusServiceUnavailable, err.Error())
return
}
authx.WriteError(w, http.StatusBadRequest, err.Error())
return
}
n := 0
if res != nil {
n = len(res.Tables)
}
logSyncReq(r, "schema/fingerprint", channelID, "", "", "ok", "", dur, reqID)
auditSync(svcCtx, r, "dbsync.schema_fingerprint", map[string]any{
"channel_id": channelID, "n": n, "ms": dur.Milliseconds(), "req_id": reqID,
})
httpx.OkJson(w, map[string]any{"success": true, "result": res})
}
}
func uniqueStringSlice(a, b []string) []string {
seen := map[string]struct{}{}
var out []string
for _, xs := range [][]string{a, b} {
for _, s := range xs {
if s == "" {
continue
}
if _, ok := seen[s]; ok {
continue
}
seen[s] = struct{}{}
out = append(out, s)
}
}
return out
}