diff --git a/monkeyai/backend/internal/agentconfig/resources.go b/monkeyai/backend/internal/agentconfig/resources.go index e2aee9212..ed92d1bdc 100644 --- a/monkeyai/backend/internal/agentconfig/resources.go +++ b/monkeyai/backend/internal/agentconfig/resources.go @@ -2,7 +2,9 @@ package agentconfig import ( "context" + "errors" "fmt" + "log/slog" "net/http" "slices" "strings" @@ -297,6 +299,12 @@ func (r *Resources) list(ctx context.Context, q resource.Queryer, user, kind str return out, nil } +func rollbackCatalog(ctx context.Context, tx pgx.Tx, operation, user string) { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + slog.ErrorContext(ctx, "回滚 Agent 配置事务失败", "operation", operation, "user_id", user, "error", err) + } +} + func (r *Resources) getList(w http.ResponseWriter, req *http.Request, kind string) { page, size, err := resource.PageParams(req) if err != nil { @@ -314,7 +322,7 @@ func (r *Resources) getList(w http.ResponseWriter, req *http.Request, kind strin resource.Fail(w, err) return } - defer tx.Rollback(req.Context()) + defer func() { rollbackCatalog(req.Context(), tx, "catalog_read", u.ID) }() items, err := r.list(req.Context(), tx, u.ID, kind, filter) if err == nil { items = resource.FilterTags(items, resource.QueryTagIDs(req)) @@ -333,7 +341,7 @@ func (r *Resources) getTags(w http.ResponseWriter, req *http.Request, kind strin resource.Fail(w, err) return } - defer tx.Rollback(req.Context()) + defer func() { rollbackCatalog(req.Context(), tx, "catalog_read", u.ID) }() items, err := r.list(req.Context(), tx, u.ID, kind, resource.CatalogFilter{}) if err == nil { err = httpapi.CachedJSON(w, req, map[string]any{"tags": resource.CollectTags(items)}) @@ -365,7 +373,7 @@ func (r *Resources) getManifest(w http.ResponseWriter, req *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(req.Context()) + defer func() { rollbackCatalog(req.Context(), tx, "catalog_read", u.ID) }() c, err := r.load(req.Context(), tx, u.ID, "") if err != nil { resource.Fail(w, err) @@ -392,7 +400,7 @@ func (r *Resources) download(w http.ResponseWriter, req *http.Request, delegated resource.Fail(w, err) return } - defer tx.Rollback(context.WithoutCancel(req.Context())) + defer func() { rollbackCatalog(context.WithoutCancel(req.Context()), tx, "download", u.ID) }() c, err := r.load(req.Context(), tx, u.ID, "") if err != nil { resource.Fail(w, err) @@ -443,7 +451,7 @@ func (r *Resources) Resolve(ctx context.Context, user string, in resource.Object if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer func() { rollbackCatalog(ctx, tx, "resolve", user) }() c, err := r.load(ctx, tx, user, "") if err != nil { return nil, err diff --git a/monkeyai/backend/internal/apikey/admin.go b/monkeyai/backend/internal/apikey/admin.go index 8e6d6552b..1c0b509e6 100644 --- a/monkeyai/backend/internal/apikey/admin.go +++ b/monkeyai/backend/internal/apikey/admin.go @@ -2,6 +2,7 @@ package apikey import ( "errors" + "log/slog" "net/http" "github.com/go-chi/chi/v5" @@ -11,6 +12,7 @@ func (s *Service) RegisterAdmin(router chi.Router) { router.Get("/api-keys", func(w http.ResponseWriter, r *http.Request) { keys, err := s.AdminList(r.Context(), r.URL.Query().Get("user_id")) if err != nil { + slog.ErrorContext(r.Context(), "读取管理员调用密钥列表失败", "error", err) keyError(w, http.StatusInternalServerError, "读取调用密钥失败") return } @@ -21,6 +23,8 @@ func (s *Service) RegisterAdmin(router chi.Router) { status := http.StatusInternalServerError if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "撤销调用密钥失败", "key_id", chi.URLParam(r, "keyID"), "error", err) } keyError(w, status, "调用密钥不存在") return diff --git a/monkeyai/backend/internal/apikey/agent.go b/monkeyai/backend/internal/apikey/agent.go index edc57d301..9786c2be2 100644 --- a/monkeyai/backend/internal/apikey/agent.go +++ b/monkeyai/backend/internal/apikey/agent.go @@ -3,6 +3,7 @@ package apikey import ( "encoding/json" "errors" + "log/slog" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity" @@ -14,6 +15,7 @@ func (s *Service) RegisterAgent(router chi.Router) { user, _ := identity.UserFromContext(r.Context()) keys, err := s.ListByUser(r.Context(), user.ID) if err != nil { + slog.ErrorContext(r.Context(), "读取用户调用密钥失败", "user_id", user.ID, "error", err) keyError(w, http.StatusInternalServerError, "读取调用密钥失败") return } @@ -47,6 +49,8 @@ func (s *Service) revoke(w http.ResponseWriter, r *http.Request) { status := http.StatusInternalServerError if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "操作用户调用密钥失败", "key_id", chi.URLParam(r, "keyID"), "user_id", user.ID, "error", err) } keyError(w, status, "调用密钥不存在") return @@ -61,6 +65,8 @@ func (s *Service) rotate(w http.ResponseWriter, r *http.Request) { status := http.StatusInternalServerError if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "操作用户调用密钥失败", "key_id", chi.URLParam(r, "keyID"), "user_id", user.ID, "error", err) } keyError(w, status, "轮换调用密钥失败") return @@ -72,7 +78,9 @@ func keyJSON(w http.ResponseWriter, status int, value any) { w.Header().Set("Cache-Control", "private, no-store") w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(value) + if err := json.NewEncoder(w).Encode(value); err != nil { + slog.Error("写入调用密钥 HTTP 响应失败", "status", status, "error", err) + } } func keyError(w http.ResponseWriter, status int, message string) { diff --git a/monkeyai/backend/internal/apikey/postgres.go b/monkeyai/backend/internal/apikey/postgres.go index 74a2b83eb..e369f835a 100644 --- a/monkeyai/backend/internal/apikey/postgres.go +++ b/monkeyai/backend/internal/apikey/postgres.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "log/slog" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/apikey/sqlc" @@ -82,6 +83,8 @@ func (p *Postgres) Authenticate(ctx context.Context, keyHash, scope string) (str if err != nil { return "", fmt.Errorf("验证调用密钥: %w", err) } - _, _ = sqlc.New(p.pool).TouchKey(ctx, id) + if _, err := sqlc.New(p.pool).TouchKey(ctx, id); err != nil { + slog.ErrorContext(ctx, "更新调用密钥使用时间失败", "key_id", id, "error", err) + } return userID, nil } diff --git a/monkeyai/backend/internal/apikey/service.go b/monkeyai/backend/internal/apikey/service.go index 289e1ebaa..9e65d0c09 100644 --- a/monkeyai/backend/internal/apikey/service.go +++ b/monkeyai/backend/internal/apikey/service.go @@ -8,6 +8,7 @@ import ( "encoding/hex" "errors" "fmt" + "log/slog" "slices" "strings" "time" @@ -76,6 +77,7 @@ func (s *Service) Create(ctx context.Context, userID string, input CreateInput) } stored, err := s.store.Create(ctx, key, hash(raw)) if err != nil { + slog.ErrorContext(ctx, "保存调用密钥失败", "user_id", userID, "error", err) return CreatedKey{}, err } return CreatedKey{Key: stored, APIKey: raw}, nil @@ -108,7 +110,9 @@ func (s *Service) Rotate(ctx context.Context, userID, id string) (CreatedKey, er return CreatedKey{}, err } if err := s.store.Revoke(ctx, id, userID); err != nil { - _ = s.store.Revoke(ctx, created.ID, userID) + if rollbackErr := s.store.Revoke(ctx, created.ID, userID); rollbackErr != nil { + slog.ErrorContext(ctx, "轮换调用密钥失败后撤销新密钥失败", "key_id", created.ID, "user_id", userID, "error", rollbackErr) + } return CreatedKey{}, err } return created, nil @@ -136,6 +140,9 @@ func (s *Service) Authenticate(ctx context.Context, raw, scope string) (string, } userID, err := s.store.Authenticate(ctx, hash(raw), scope) if err != nil { + if !errors.Is(err, ErrInvalidKey) { + slog.ErrorContext(ctx, "验证调用密钥失败", "scope", scope, "error", err) + } return "", ErrInvalidKey } return userID, nil diff --git a/monkeyai/backend/internal/apikey/service_test.go b/monkeyai/backend/internal/apikey/service_test.go index 37f6b67fd..d31cb4103 100644 --- a/monkeyai/backend/internal/apikey/service_test.go +++ b/monkeyai/backend/internal/apikey/service_test.go @@ -1,22 +1,27 @@ package apikey import ( + "bytes" "context" "errors" + "fmt" + "log/slog" "strings" "testing" "time" ) type storeStub struct { - keys []Key - hash string - userID string - authErr error + keys []Key + hash string + userID string + authErr error + revokeErr map[string]error + revoked []string } func (s *storeStub) Create(_ context.Context, key Key, hash string) (Key, error) { - key.ID = "key-1" + key.ID = fmt.Sprintf("key-%d", len(s.keys)+1) key.CreatedAt = time.Date(2026, 9, 6, 0, 0, 0, 0, time.UTC) s.keys = append(s.keys, key) s.hash = hash @@ -25,7 +30,10 @@ func (s *storeStub) Create(_ context.Context, key Key, hash string) (Key, error) func (s *storeStub) ListByUser(context.Context, string) ([]Key, error) { return s.keys, nil } func (s *storeStub) List(context.Context, string) ([]Key, error) { return s.keys, nil } -func (s *storeStub) Revoke(context.Context, string, string) error { return nil } +func (s *storeStub) Revoke(_ context.Context, id, _ string) error { + s.revoked = append(s.revoked, id) + return s.revokeErr[id] +} func (s *storeStub) Authenticate(context.Context, string, string) (string, error) { return s.userID, s.authErr } @@ -79,3 +87,21 @@ func TestAuthenticateHidesStoreErrors(t *testing.T) { t.Fatalf("error = %v", err) } } + +func TestRotateLogsFailedCompensationWithoutSecret(t *testing.T) { + var output bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&output, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + store := &storeStub{ + keys: []Key{{ID: "key-1", Name: "work", Scopes: []string{ScopeModelInvoke}, ExpiresAt: time.Now().Add(30 * 24 * time.Hour)}}, + revokeErr: map[string]error{"key-1": errors.New("old revoke failed"), "key-2": errors.New("cleanup failed")}, + } + _, err := NewService(store).Rotate(t.Context(), "user-1", "key-1") + if err == nil || len(store.revoked) != 2 || store.revoked[0] != "key-1" || store.revoked[1] != "key-2" { + t.Fatalf("撤销失败后应尝试回收新密钥: %v, %v", store.revoked, err) + } + if !strings.Contains(output.String(), "cleanup failed") || !strings.Contains(output.String(), "key-2") || strings.Contains(output.String(), store.hash) { + t.Fatalf("补偿日志缺少错误或包含密钥摘要: %s", output.String()) + } +} diff --git a/monkeyai/backend/internal/audit/middleware.go b/monkeyai/backend/internal/audit/middleware.go index 95fd969d0..0f77ca239 100644 --- a/monkeyai/backend/internal/audit/middleware.go +++ b/monkeyai/backend/internal/audit/middleware.go @@ -42,7 +42,9 @@ func (b *capture) Write(p []byte) (int, error) { if n > bodyLimit-b.Len() { b.overflow = true } - _, _ = b.Buffer.Write(p[:min(n, bodyLimit-b.Len())]) + if _, err := b.Buffer.Write(p[:min(n, bodyLimit-b.Len())]); err != nil { + return 0, err + } return n, nil } @@ -67,13 +69,17 @@ func (s *Service) Middleware(actor func(*http.Request) Actor) func(http.Handler) } var random [16]byte if _, err := rand.Read(random[:]); err != nil { - failure(w, 500, "server_error", "初始化操作审计失败") + s.logger.ErrorContext(r.Context(), "初始化操作审计失败", "error", err) + failure(r.Context(), w, 500, "server_error", "初始化操作审计失败") return } state := &requestState{id: hex.EncodeToString(random[:]), actor: user, at: time.Now().UTC(), ip: clientIP(r), agent: clean(r.UserAgent(), 512)} r = r.WithContext(context.WithValue(r.Context(), requestKey{}, state)) var body, response capture - media, _, _ := mime.ParseMediaType(r.Header.Get("Content-Type")) + media, _, mediaErr := mime.ParseMediaType(r.Header.Get("Content-Type")) + if mediaErr != nil && r.Header.Get("Content-Type") != "" { + s.logger.WarnContext(r.Context(), "操作审计请求类型无法解析", "error", mediaErr) + } if r.Body != nil && media == "application/json" { r.Body = bodyReader{Reader: io.TeeReader(r.Body, &body), Closer: r.Body} } @@ -91,7 +97,7 @@ func (s *Service) Middleware(actor func(*http.Request) Actor) func(http.Handler) ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 5*time.Second) defer cancel() if err := s.complete(ctx, r, state, &body, &response, media, status); err != nil { - s.logger.Error("操作审计写入失败", "request_id", state.id, "method", r.Method, "status", status) + s.logger.ErrorContext(ctx, "操作审计写入失败", "request_id", state.id, "method", r.Method, "status", status, "error", err) } if panicked != nil { panic(panicked) diff --git a/monkeyai/backend/internal/audit/service.go b/monkeyai/backend/internal/audit/service.go index d60f98205..e3a656004 100644 --- a/monkeyai/backend/internal/audit/service.go +++ b/monkeyai/backend/internal/audit/service.go @@ -14,6 +14,7 @@ import ( "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/audit/sqlc" "github.com/go-chi/chi/v5" + "github.com/go-chi/chi/v5/middleware" "github.com/jackc/pgx/v5/pgtype" "github.com/jackc/pgx/v5/pgxpool" ) @@ -125,14 +126,15 @@ func (s *Service) list(w http.ResponseWriter, r *http.Request) { w.Header().Set("Cache-Control", "private, no-store") in, page, err := filters(r) if err != nil { - failure(w, 400, "invalid_request", err.Error()) + failure(r.Context(), w, 400, "invalid_request", err.Error()) return } ctx, cancel := context.WithTimeout(r.Context(), 15*time.Second) defer cancel() data, err := sqlc.New(s.pool).Page(ctx, in) if err != nil { - failure(w, 500, "server_error", "读取操作审计失败") + s.logger.ErrorContext(r.Context(), "读取操作审计失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + failure(r.Context(), w, 500, "server_error", "读取操作审计失败") return } var out struct { @@ -144,7 +146,8 @@ func (s *Service) list(w http.ResponseWriter, r *http.Request) { decoder := json.NewDecoder(bytes.NewReader(data)) decoder.UseNumber() if err = decoder.Decode(&out); err != nil { - failure(w, 500, "server_error", "读取操作审计失败") + s.logger.ErrorContext(r.Context(), "读取操作审计失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + failure(r.Context(), w, 500, "server_error", "读取操作审计失败") return } for _, item := range out.Items { @@ -152,11 +155,15 @@ func (s *Service) list(w http.ResponseWriter, r *http.Request) { } out.Page, out.PageSize = page, in.PageSize w.Header().Set("Content-Type", "application/json; charset=utf-8") - _ = json.NewEncoder(w).Encode(out) + if err := json.NewEncoder(w).Encode(out); err != nil { + s.logger.ErrorContext(r.Context(), "操作审计响应写入失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + } } -func failure(w http.ResponseWriter, status int, code, message string) { +func failure(ctx context.Context, w http.ResponseWriter, status int, code, message string) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(map[string]any{"error": map[string]string{"code": code, "message": message}}) + if err := json.NewEncoder(w).Encode(map[string]any{"error": map[string]string{"code": code, "message": message}}); err != nil { + slog.ErrorContext(ctx, "操作审计错误响应写入失败", "request_id", middleware.GetReqID(ctx), "status", status, "code", code, "error", err) + } } diff --git a/monkeyai/backend/internal/billing/admin.go b/monkeyai/backend/internal/billing/admin.go index 3fcf07118..890cbfd00 100644 --- a/monkeyai/backend/internal/billing/admin.go +++ b/monkeyai/backend/internal/billing/admin.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "net/http" "slices" "strconv" @@ -101,7 +102,7 @@ func (s *Service) saveSettings(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer rollback(r.Context(), tx, "save_settings", "") p, err := s.policy(r.Context(), tx, true) if err != nil { resource.Fail(w, err) @@ -158,7 +159,11 @@ func (s *Service) saveSettings(w http.ResponseWriter, r *http.Request) { } } p.Revision++ - raw, _ := json.Marshal(p) + raw, err := json.Marshal(p) + if err != nil { + resource.Fail(w, fmt.Errorf("序列化计费策略: %w", err)) + return + } u, _ := identity.UserFromContext(r.Context()) _, err = sqlc.New(tx).SavePolicy(r.Context(), sqlc.SavePolicyParams{Value: raw, Revision: int64(p.Revision), UpdatedByUserID: u.ID}) if err == nil { @@ -237,7 +242,7 @@ func (s *Service) saveQuotas(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer rollback(r.Context(), tx, "save_quotas", "") p, err := s.policy(r.Context(), tx, true) if err != nil { resource.Fail(w, err) @@ -387,7 +392,7 @@ func (s *Service) resetQuotas(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "reset_credits", "") queries := sqlc.New(tx) if _, err = queries.LockGroupsForReset(ctx); err != nil { resource.Fail(w, err) @@ -486,7 +491,11 @@ func (s *Service) account(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - external, _ := sqlc.New(s.pool).WalletUser(r.Context(), user) + external, err := sqlc.New(s.pool).WalletUser(r.Context(), user) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + resource.Fail(w, fmt.Errorf("查询用户 %s 钱包绑定: %w", user, err)) + return + } wallet, err := s.wallet(r.Context(), s.pool) if err != nil { @@ -528,7 +537,7 @@ func (s *Service) adjust(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "adjust_balance", "") p, err := s.policy(ctx, tx, false) if err != nil { resource.Fail(w, err) @@ -547,7 +556,12 @@ func (s *Service) adjust(w http.ResponseWriter, r *http.Request) { } if e == nil { - if oldAmount != in.Delta.String() && amountText(oldAmount) != in.Delta || oldReason != in.Reason { + previous, parseErr := ParseAmount(oldAmount) + if parseErr != nil { + resource.Fail(w, fmt.Errorf("解析账户 %s 的历史调整金额: %w", a.ID, parseErr)) + return + } + if previous != in.Delta || oldReason != in.Reason { resource.Fail(w, fail(409, "idempotency_conflict", "该账户版本已提交过不同的调整,请刷新后重试")) return } @@ -594,15 +608,15 @@ func (s *Service) adjust(w http.ResponseWriter, r *http.Request) { s.account(w, r) } func pageParams(r *http.Request) (int, int) { - page, _ := strconv.Atoi(r.URL.Query().Get("page")) - size, _ := strconv.Atoi(r.URL.Query().Get("page_size")) - if page < 1 { + page, pageErr := strconv.Atoi(r.URL.Query().Get("page")) + size, sizeErr := strconv.Atoi(r.URL.Query().Get("page_size")) + if pageErr != nil || page < 1 { page = 1 } if page > 100000 { page = 100000 } - if size < 1 || size > 100 { + if sizeErr != nil || size < 1 || size > 100 { size = 20 } return page, size @@ -628,7 +642,7 @@ func (s *Service) refund(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "refund", id) var account, mode, status, amountText, category, item string var record sqlc.LockRefundRow record, err = sqlc.New(tx).LockRefund(ctx, id) diff --git a/monkeyai/backend/internal/billing/agent.go b/monkeyai/backend/internal/billing/agent.go index 1801ae5ce..294f30133 100644 --- a/monkeyai/backend/internal/billing/agent.go +++ b/monkeyai/backend/internal/billing/agent.go @@ -41,7 +41,7 @@ func (s *Service) agentEntries(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "list_user_entries", "") queries := sqlc.New(tx) total, err := queries.CountUserEntries(ctx, sqlc.CountUserEntriesParams{ UserID: user.ID, diff --git a/monkeyai/backend/internal/billing/agent_test.go b/monkeyai/backend/internal/billing/agent_test.go index 93b257a0b..24c73c7b9 100644 --- a/monkeyai/backend/internal/billing/agent_test.go +++ b/monkeyai/backend/internal/billing/agent_test.go @@ -109,7 +109,7 @@ func TestAgentBilling(t *testing.T) { if err != nil { t.Fatal(err) } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "test_agent_entries", "") for _, entry := range []struct{ account, kind, category, item, delta, mode string }{ {a.String("id"), "refund", "model", "模型退款", "0.48", "local"}, {a.String("id"), "charge", "tool", "工具调用", "-2.5", "remote"}, diff --git a/monkeyai/backend/internal/billing/amount.go b/monkeyai/backend/internal/billing/amount.go index 40e2004b9..c07d30483 100644 --- a/monkeyai/backend/internal/billing/amount.go +++ b/monkeyai/backend/internal/billing/amount.go @@ -106,5 +106,4 @@ func boolInt(b bool) int64 { } return 0 } -func amountText(s string) Amount { a, _ := ParseAmount(s); return a } -func stringInt(n int64) string { return strconv.FormatInt(n, 10) } +func stringInt(n int64) string { return strconv.FormatInt(n, 10) } diff --git a/monkeyai/backend/internal/billing/connection.go b/monkeyai/backend/internal/billing/connection.go index 71c0880cf..742010dcb 100644 --- a/monkeyai/backend/internal/billing/connection.go +++ b/monkeyai/backend/internal/billing/connection.go @@ -9,6 +9,7 @@ import ( "encoding/pem" "errors" "fmt" + "log/slog" "net" "net/url" "os" @@ -54,7 +55,13 @@ func WalletFromEnv() (*Wallet, error) { return nil, resource.Invalid("请配置百智云服务 URL") } } - cfg.AppID, _ = strconv.Atoi(os.Getenv("BAIZHIYUN_APP_ID")) + if raw := os.Getenv("BAIZHIYUN_APP_ID"); raw != "" { + appID, err := strconv.Atoi(raw) + if err != nil { + return nil, resource.Invalid("BAIZHIYUN_APP_ID 必须为整数") + } + cfg.AppID = appID + } dir := os.Getenv("MONKEYAI_WALLET_CERT_DIR") if dir == "" { return nil, errors.New("远程计费需配置 MONKEYAI_WALLET_CERT_DIR") @@ -62,7 +69,7 @@ func WalletFromEnv() (*Wallet, error) { for name, target := range map[string]*string{"app.crt": &cfg.Certificate, "app.key": &cfg.PrivateKey, "ca.crt": &cfg.CACertificate} { data, err := os.ReadFile(filepath.Join(dir, name)) if err != nil { - return nil, fmt.Errorf("读取钱包证书文件 %s 失败", name) + return nil, fmt.Errorf("读取钱包证书文件 %s 失败: %w", name, err) } *target = string(data) } @@ -133,12 +140,16 @@ func newWallet(cfg WalletConfig) (*Wallet, error) { // SDK 仅接收文件路径,构造完成后证书已加载到内存。 dir, err := os.MkdirTemp("", "monkeyai-wallet-") if err != nil { - return nil, errors.New("创建钱包证书临时目录失败") + return nil, fmt.Errorf("创建钱包证书临时目录失败: %w", err) } - defer os.RemoveAll(dir) + defer func() { + if err := os.RemoveAll(dir); err != nil { + slog.Error("清理钱包证书临时目录失败", "operation", "remove_wallet_temp_dir", "error", err) + } + }() for name, content := range map[string]string{"app.crt": cfg.Certificate, "app.key": cfg.PrivateKey, "ca.crt": cfg.CACertificate} { if err = os.WriteFile(filepath.Join(dir, name), []byte(content), 0600); err != nil { - return nil, errors.New("写入钱包证书临时文件失败") + return nil, fmt.Errorf("写入钱包证书临时文件 %s 失败: %w", name, err) } } openURL, walletURL := walletEndpoints(cfg.BaseURL) @@ -241,7 +252,10 @@ func (s *Service) saveWallet(ctx context.Context, q resource.Queryer, in map[str if pending { return WalletInfo{}, WalletInfo{}, fail(409, "wallet_transactions_pending", "存在未完成的远程交易,请处理后再切换服务 URL 或应用 ID;同一应用可以更新证书") } - raw, _ := json.Marshal(next) + raw, err := json.Marshal(next) + if err != nil { + return WalletInfo{}, WalletInfo{}, fmt.Errorf("序列化钱包配置失败: %w", err) + } _, err = sqlc.New(q).SaveWalletConfig(ctx, raw) return previous.info(), wallet.info(), err } diff --git a/monkeyai/backend/internal/billing/connection_test.go b/monkeyai/backend/internal/billing/connection_test.go index 0e65ffe4d..44f0f823f 100644 --- a/monkeyai/backend/internal/billing/connection_test.go +++ b/monkeyai/backend/internal/billing/connection_test.go @@ -374,3 +374,13 @@ func TestWalletPublicURLs(t *testing.T) { } } } + +func TestWalletFromEnvInvalidAppIDDoesNotExposeValue(t *testing.T) { + t.Setenv("BAIZHIYUN_ENV", "dev") + t.Setenv("BAIZHIYUN_BASE_URL", "") + t.Setenv("BAIZHIYUN_APP_ID", "private-credential") + _, err := WalletFromEnv() + if err == nil || !strings.Contains(err.Error(), "BAIZHIYUN_APP_ID") || strings.Contains(err.Error(), "private-credential") { + t.Fatalf("应用 ID 配置错误未安全报告: %v", err) + } +} diff --git a/monkeyai/backend/internal/billing/query.go b/monkeyai/backend/internal/billing/query.go index 57ea5723a..ce0f2f21c 100644 --- a/monkeyai/backend/internal/billing/query.go +++ b/monkeyai/backend/internal/billing/query.go @@ -2,6 +2,7 @@ package billing import ( "encoding/json" + "fmt" "net/http" "strings" "time" @@ -56,7 +57,7 @@ func (s *Service) entries(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer rollback(r.Context(), tx, "list_entries", "") var total int64 total, err = sqlc.New(tx).CountEntries(r.Context(), filter) if err != nil { @@ -186,7 +187,7 @@ func (s *Service) resolve(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer rollback(r.Context(), tx, "resolve_transaction", id) var state, mode string var record sqlc.LockTransactionStatusRow record, err = sqlc.New(tx).LockTransactionStatus(r.Context(), id) @@ -244,7 +245,11 @@ func (s *Service) resolve(w http.ResponseWriter, r *http.Request) { } u, _ := identity.UserFromContext(r.Context()) in.Usage.Known = true - body, _ := json.Marshal(in) + body, err := json.Marshal(in) + if err != nil { + resource.Fail(w, fmt.Errorf("序列化交易 %s 核查记录: %w", id, err)) + return + } err = audit(r.Context(), tx, u.ID, "resolve_transaction", id, json.RawMessage(body)) if err == nil { diff --git a/monkeyai/backend/internal/billing/reconcile.go b/monkeyai/backend/internal/billing/reconcile.go index 335d46a63..f4c3483c2 100644 --- a/monkeyai/backend/internal/billing/reconcile.go +++ b/monkeyai/backend/internal/billing/reconcile.go @@ -82,7 +82,10 @@ func (s *Service) reconcileUsage(ctx context.Context, row sqlc.TransactionsToRec releaseCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if _, releaseErr := sqlc.New(conn).ReleaseReconciliationLock(releaseCtx, row.ID); releaseErr != nil { - conn.Conn().Close(releaseCtx) + slog.ErrorContext(releaseCtx, "释放对账锁失败", "transaction_id", row.ID, "operation", "release_reconciliation_lock", "error", releaseErr) + if closeErr := conn.Conn().Close(releaseCtx); closeErr != nil { + slog.ErrorContext(releaseCtx, "关闭对账连接失败", "transaction_id", row.ID, "operation", "close_reconciliation_connection", "error", closeErr) + } } }() diff --git a/monkeyai/backend/internal/billing/service.go b/monkeyai/backend/internal/billing/service.go index 4cd31e170..76246b7e3 100644 --- a/monkeyai/backend/internal/billing/service.go +++ b/monkeyai/backend/internal/billing/service.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "sync" "time" @@ -46,7 +47,11 @@ func defaultPolicy() Policy { return Policy{RootCredits: 10000 * Amount(scale), Input: 100 * Amount(scale), Cached: 20 * Amount(scale), Output: 400 * Amount(scale), Cycle: "weekly", Mode: "local", Enabled: true} } func (p Policy) period(now time.Time) (time.Time, time.Time) { - zone, _ := time.LoadLocation("Asia/Shanghai") + zone, err := time.LoadLocation("Asia/Shanghai") + if err != nil { + slog.Error("加载计费时区失败,使用固定东八区", "operation", "load_billing_timezone", "error", err) + zone = time.FixedZone("CST", 8*60*60) + } n := now.In(zone) cycle := p.Cycle anchor := p.CycleAnchor @@ -124,8 +129,11 @@ func (s *Service) WithUsageReconciler(r UsageReconciler) *Service { } func (s *Service) Initialize(ctx context.Context) error { p := defaultPolicy() - b, _ := json.Marshal(p) - _, err := sqlc.New(s.pool).InitializePolicy(ctx, b) + b, err := json.Marshal(p) + if err != nil { + return fmt.Errorf("序列化默认计费策略: %w", err) + } + _, err = sqlc.New(s.pool).InitializePolicy(ctx, b) if err != nil { return err } @@ -255,12 +263,18 @@ func (s *Service) ensureAccount(ctx context.Context, tx pgx.Tx, user string, p P } return accountRow(ctx, tx, user, start) } +func rollback(ctx context.Context, tx pgx.Tx, operation, transactionID string) { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) { + slog.ErrorContext(ctx, "回滚计费事务失败", "operation", operation, "transaction_id", transactionID, "error", err) + } +} + func (s *Service) Account(ctx context.Context, user string) (Account, error) { tx, err := s.pool.Begin(ctx) if err != nil { return Account{}, err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "account", "") p, err := s.policy(ctx, tx, false) if err != nil { return Account{}, err diff --git a/monkeyai/backend/internal/billing/service_test.go b/monkeyai/backend/internal/billing/service_test.go index a3990d738..d101531ed 100644 --- a/monkeyai/backend/internal/billing/service_test.go +++ b/monkeyai/backend/internal/billing/service_test.go @@ -551,7 +551,7 @@ func TestQuotaChangePreservesUnopenedAccount(t *testing.T) { if err != nil { t.Fatal(err) } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "test_preserve_accounts", "") if err = s.PreserveAccounts(ctx, tx); err != nil { t.Fatal(err) } @@ -566,3 +566,11 @@ func TestQuotaChangePreservesUnopenedAccount(t *testing.T) { t.Fatal("首次访问不应提前使用下周期额度", a, err) } } + +func amountText(s string) Amount { + a, err := ParseAmount(s) + if err != nil { + panic(err) + } + return a +} diff --git a/monkeyai/backend/internal/billing/transaction.go b/monkeyai/backend/internal/billing/transaction.go index 7a83c92ad..15da0e88b 100644 --- a/monkeyai/backend/internal/billing/transaction.go +++ b/monkeyai/backend/internal/billing/transaction.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "log/slog" "math" "time" @@ -45,7 +46,7 @@ func (s *Service) Begin(ctx context.Context, r Request) (Reservation, error) { if err != nil { return Reservation{}, err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "begin", "") p, err := s.policy(ctx, tx, false) if err != nil { return Reservation{}, err @@ -237,7 +238,10 @@ func (s *Service) Begin(ctx context.Context, r Request) (Reservation, error) { return Reservation{}, err } } - snapshot, _ := json.Marshal(price) + snapshot, err := json.Marshal(price) + if err != nil { + return Reservation{}, fmt.Errorf("序列化计费价格快照: %w", err) + } id := resource.ID() _, err = sqlc.New(tx).FreezeBalance(ctx, sqlc.FreezeBalanceParams{ID: a.ID, Frozen: reserve.String()}) @@ -313,7 +317,7 @@ func (s *Service) Finish(ctx context.Context, id string, u Usage) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "finish", id) var category, mode, status, reserveText string var raw []byte var record sqlc.LockUsageRow @@ -368,7 +372,10 @@ func (s *Service) Finish(ctx context.Context, id string, u Usage) error { state = "unknown" code = "reservation_exceeded" } - usage, _ := json.Marshal(u) + usage, err := json.Marshal(u) + if err != nil { + return fmt.Errorf("序列化交易 %s 的用量: %w", id, err) + } _, err = sqlc.New(tx).SaveUsage(ctx, sqlc.SaveUsageParams{ ID: id, Status: state, @@ -421,7 +428,10 @@ func (s *Service) Settle(ctx context.Context, id string) error { c, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if _, e := sqlc.New(conn).ReleaseSettlementLock(c, id); e != nil { - conn.Conn().Close(c) + slog.ErrorContext(c, "释放结算锁失败", "transaction_id", id, "operation", "release_settlement_lock", "error", e) + if closeErr := conn.Conn().Close(c); closeErr != nil { + slog.ErrorContext(c, "关闭结算连接失败", "transaction_id", id, "operation", "close_settlement_connection", "error", closeErr) + } } }() var mode, status string @@ -444,7 +454,7 @@ func (s *Service) Settle(ctx context.Context, id string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "settle", id) var account, item, category, amountText, reserveText string var settlement sqlc.LockSettlementRow settlement, err = sqlc.New(tx).LockSettlement(ctx, id) @@ -498,7 +508,10 @@ func (s *Service) recover(ctx context.Context) error { e := s.Settle(c, id) cancel() if e != nil { - _, _ = sqlc.New(s.pool).ScheduleRetry(ctx, id) + slog.WarnContext(ctx, "恢复交易结算失败", "transaction_id", id, "operation", "settle", "error", e) + if _, retryErr := sqlc.New(s.pool).ScheduleRetry(ctx, id); retryErr != nil { + return fmt.Errorf("安排交易 %s 结算重试: %w", id, retryErr) + } } } return s.reconcileUnknown(ctx) diff --git a/monkeyai/backend/internal/billing/wallet.go b/monkeyai/backend/internal/billing/wallet.go index d0c0015fb..59c4423a8 100644 --- a/monkeyai/backend/internal/billing/wallet.go +++ b/monkeyai/backend/internal/billing/wallet.go @@ -5,6 +5,7 @@ import ( "crypto/rand" "errors" "fmt" + "log/slog" "math/big" "time" @@ -93,7 +94,7 @@ func (s *Service) reserveRemote(ctx context.Context, id string) error { if e != nil { return e } - defer tx.Rollback(finalCtx) + defer rollback(finalCtx, tx, "mark_wallet_reservation_failed", id) state := "unknown" if definite { state = "rejected" @@ -125,7 +126,7 @@ func (s *Service) reserveRemote(ctx context.Context, id string) error { if err != nil { return err } - defer tx.Rollback(finalCtx) + defer rollback(finalCtx, tx, "mark_wallet_reserved", id) _, err = sqlc.New(tx).MarkWalletReserved(finalCtx, id) if err != nil { return err @@ -177,8 +178,12 @@ func (s *Service) confirmRemote(ctx context.Context, conn *pgxpool.Conn, id stri err = wallet.Client.ConfirmBillingCharge(ctx, &opensdk.ConfirmBillingChargeReq{BizID: biz, UserID: user, TeamSlug: team, Status: "success", ActualAmountCreditCents: quotaAmount(amount, false), Subject: item}) if err != nil { code, trace := walletFailure(err) - _, _ = sqlc.New(conn).SetWalletError(ctx, sqlc.SetWalletErrorParams{TransactionID: id, ErrorCode: code, TraceID: trace}) - _, _ = sqlc.New(conn).SetTransactionError(ctx, sqlc.SetTransactionErrorParams{ID: id, ErrorCode: code}) + if _, markErr := sqlc.New(conn).SetWalletError(ctx, sqlc.SetWalletErrorParams{TransactionID: id, ErrorCode: code, TraceID: trace}); markErr != nil { + slog.ErrorContext(ctx, "记录钱包确认错误失败", "transaction_id", id, "operation", "set_wallet_error", "error", markErr) + } + if _, markErr := sqlc.New(conn).SetTransactionError(ctx, sqlc.SetTransactionErrorParams{ID: id, ErrorCode: code}); markErr != nil { + slog.ErrorContext(ctx, "记录交易确认错误失败", "transaction_id", id, "operation", "set_transaction_error", "error", markErr) + } return fail(503, code, "百智云确认待重试") } _, err = sqlc.New(conn).MarkWalletConfirmed(ctx, id) @@ -215,7 +220,7 @@ func (s *Service) BindWallet(ctx context.Context, actor, user, external string) if err != nil { return err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, "bind_wallet", "") _, err = sqlc.New(tx).LockUser(ctx, user) if err != nil { return err diff --git a/monkeyai/backend/internal/endpoint/connection.go b/monkeyai/backend/internal/endpoint/connection.go index f65c0b03d..fa38a5851 100644 --- a/monkeyai/backend/internal/endpoint/connection.go +++ b/monkeyai/backend/internal/endpoint/connection.go @@ -3,8 +3,13 @@ package endpoint import ( "context" "encoding/json" + "errors" + "fmt" + "io" + "net" "sync" "sync/atomic" + "syscall" "time" "unicode/utf8" @@ -71,7 +76,12 @@ func (c *connection) reject(code, id string) { if !validID(id) { id = "" } - data := errorFrame(code, id) + data, err := errorFrame(code, id) + if err != nil { + c.service.logger.Error("序列化端点错误帧失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "reject", "error", err) + c.stop(1011) + return + } c.service.stats.rejected.Add(1) c.mu.Lock() defer c.mu.Unlock() @@ -112,10 +122,27 @@ func (c *connection) pop(high bool) ([]byte, bool) { } return nil, false } +func (c *connection) expected(err error) bool { + return c.ctx.Err() != nil || errors.Is(err, context.Canceled) || errors.Is(err, io.EOF) || + errors.Is(err, net.ErrClosed) || errors.Is(err, syscall.EPIPE) || + errors.Is(err, syscall.ECONNRESET) || websocket.CloseStatus(err) != -1 +} +func (c *connection) sendError(code, reply string) { + data, err := errorFrame(code, reply) + if err != nil { + c.service.logger.Error("序列化端点错误帧失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "send_error", "error", err) + c.stop(1011) + return + } + c.write(data) +} func (c *connection) write(data []byte) bool { ctx, cancel := context.WithTimeout(c.ctx, c.service.timeWrite) defer cancel() if err := c.ws.Write(ctx, websocket.MessageText, data); err != nil { + if !c.dead.Load() && !c.expected(err) { + c.service.logger.Warn("端点消息写入失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "ws_write", "error", err) + } c.stop(1013) return false } @@ -127,7 +154,9 @@ func (c *connection) writer() { if code == 0 { code = 1000 } - _ = c.ws.Close(websocket.StatusCode(code), "") + if err := c.ws.Close(websocket.StatusCode(code), ""); err != nil && !c.expected(err) { + c.service.logger.Warn("端点 WebSocket 关闭失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "ws_close", "close_code", code, "error_type", fmt.Sprintf("%T", err)) + } c.cancel() }() select { @@ -140,7 +169,12 @@ func (c *connection) writer() { return } welcome := map[string]any{"type": "welcome", "protocol_version": 1, "server_time": time.Now().UnixMilli(), "heartbeat": map[string]int64{"interval_ms": c.service.pingInterval.Milliseconds(), "timeout_ms": c.service.timePong.Milliseconds()}, "limits": map[string]int{"max_frame_bytes": maxMessage, "max_endpoints": 20}} - data, _ := json.Marshal(welcome) + data, err := json.Marshal(welcome) + if err != nil { + c.service.logger.Error("序列化端点欢迎消息失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "welcome", "error", err) + c.stop(1011) + return + } if !c.write(data) { return } @@ -189,6 +223,9 @@ func (c *connection) reader() { for { kind, data, err := c.ws.Read(c.ctx) if err != nil { + if !c.dead.Load() && !c.expected(err) { + c.service.logger.Warn("端点消息读取失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "ws_read", "error", err) + } c.stop(1000) return } @@ -202,8 +239,13 @@ func (c *connection) reader() { } var envelope map[string]json.RawMessage var messageType string - _ = json.Unmarshal(data, &envelope) - _ = json.Unmarshal(envelope["type"], &messageType) + if err := json.Unmarshal(data, &envelope); err == nil { + if raw, ok := envelope["type"]; ok { + if err := json.Unmarshal(raw, &messageType); err != nil { + messageType = "" + } + } + } if messageType == "hello" { c.stop(1002) return @@ -237,6 +279,9 @@ func (c *connection) ping() { err := c.ws.Ping(ctx) cancel() if err != nil { + if !c.dead.Load() && !c.expected(err) { + c.service.logger.Warn("端点心跳失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "ws_ping", "error", err) + } c.stop(1013) return } @@ -254,6 +299,9 @@ func (c *connection) ping() { err = c.service.store.Touch(ctx, c.credential.UserID, c.machine, now) cancel() if err != nil { + if !c.dead.Load() && !c.expected(err) { + c.service.logger.Warn("端点心跳刷新失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "touch", "error", err) + } c.stop(1013) return } @@ -274,7 +322,10 @@ func (c *connection) verify() { case <-c.ctx.Done(): return case <-ticker.C: - code, _ := c.service.check(c.ctx, c.credential) + code, err := c.service.check(c.ctx, c.credential) + if err != nil && !c.expected(err) { + c.service.logger.Warn("端点凭据复核失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "verify", "error", err) + } if code != 0 { c.stop(code) return diff --git a/monkeyai/backend/internal/endpoint/endpoint_test.go b/monkeyai/backend/internal/endpoint/endpoint_test.go index 3a826c12b..e96f3853a 100644 --- a/monkeyai/backend/internal/endpoint/endpoint_test.go +++ b/monkeyai/backend/internal/endpoint/endpoint_test.go @@ -189,7 +189,11 @@ func call(t *testing.T, f *fixture, method, path, body, token string) (int, []by if err != nil { t.Fatal(err) } - defer response.Body.Close() + defer func() { + if err := response.Body.Close(); err != nil && t.Context().Err() == nil { + t.Errorf("关闭端点测试响应失败: %T", err) + } + }() data, _ := io.ReadAll(response.Body) return response.StatusCode, data } diff --git a/monkeyai/backend/internal/endpoint/handler.go b/monkeyai/backend/internal/endpoint/handler.go index d6572b0c9..28887721d 100644 --- a/monkeyai/backend/internal/endpoint/handler.go +++ b/monkeyai/backend/internal/endpoint/handler.go @@ -5,11 +5,14 @@ import ( "encoding/json" "errors" "io" + "log/slog" + "net" "net/http" "net/url" "strconv" "strings" "sync" + "syscall" "time" "unicode/utf8" @@ -30,8 +33,20 @@ func (s *Service) RegisterAgent(router chi.Router) { func respond(w http.ResponseWriter, code int, data any) { w.Header().Set("Cache-Control", "private, no-store") w.Header().Set("Content-Type", "application/json; charset=utf-8") + body, err := json.Marshal(data) + if err != nil { + slog.Error("序列化端点响应失败", "operation", "respond", "status", code, "error", err) + http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) + return + } w.WriteHeader(code) - _ = json.NewEncoder(w).Encode(data) + if _, err := w.Write(append(body, '\n')); err != nil { + if errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || errors.Is(err, syscall.EPIPE) || errors.Is(err, syscall.ECONNRESET) { + slog.Debug("端点响应连接已断开", "operation", "respond", "status", code, "error", err) + } else { + slog.Warn("端点响应写入失败", "operation", "respond", "status", code, "error", err) + } + } } func failure(w http.ResponseWriter, code int, name string) { respond(w, code, map[string]any{"error": map[string]string{"code": name, "message": map[string]string{"invalid_request": "请求参数无效", "invalid_token": "凭据无效", "endpoint_not_found": "端点不存在", "endpoint_limit_exceeded": "端点数量达到上限", "service_unavailable": "服务暂不可用", "forbidden": "请求来源不受信任", "rate_limited": "请求过于频繁"}[name]}}) @@ -66,6 +81,9 @@ func (s *Service) manage(w http.ResponseWriter, r *http.Request, fn func(context u := s.acquire(credential.UserID) defer s.release(credential.UserID, u) if err := u.enter(ctx); err != nil { + if r.Context().Err() == nil { + s.logger.Warn("端点管理获取锁失败", "user_id", credential.UserID, "operation", r.Method, "error", err) + } failure(w, 503, "service_unavailable") return } @@ -75,12 +93,18 @@ func (s *Service) manage(w http.ResponseWriter, r *http.Request, fn func(context return } if err := s.load(ctx, credential.UserID, u); err != nil { + if r.Context().Err() == nil { + s.logger.Warn("加载端点目录失败", "user_id", credential.UserID, "operation", r.Method, "error", err) + } failure(w, 503, "service_unavailable") return } result, err := fn(ctx, credential.UserID, u) if err != nil { code, name := status(err) + if name == "service_unavailable" && r.Context().Err() == nil { + s.logger.Warn("端点管理操作失败", "user_id", credential.UserID, "operation", r.Method, "error", err) + } failure(w, code, name) return } @@ -149,7 +173,7 @@ func (s *Service) update(w http.ResponseWriter, r *http.Request, action string) } else { u.endpoints[machine] = e } - s.broadcast(u) + s.broadcast(user, u) return s.view(u, e), nil }) } @@ -209,6 +233,9 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { defer func() { s.mu.Lock(); s.slots--; s.mu.Unlock() }() ws, err := websocket.Accept(w, r, &websocket.AcceptOptions{InsecureSkipVerify: true, CompressionMode: websocket.CompressionDisabled}) if err != nil { + if r.Context().Err() == nil { + s.logger.Warn("端点 WebSocket 握手失败", "user_id", credential.UserID, "operation", "ws_accept", "error", err) + } s.stats.handshakes.Add(1) return } @@ -233,19 +260,24 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { cleanup, cancel := context.WithTimeout(context.Background(), s.timeDB) defer cancel() // 清理必须完成,不能因管理操作持锁而遗留一个永久离线连接。 - _ = u.enter(context.Background()) - if u.connections[c.machine] == c { - delete(u.connections, c.machine) - if e, ok := u.endpoints[c.machine]; ok { - seen := c.seen.Load() - e.LastSeenAt = &seen - u.endpoints[c.machine] = e + if err := u.enter(context.Background()); err != nil { + s.logger.Error("清理端点连接获取锁失败", "user_id", credential.UserID, "machine_id", c.machine, "operation", "disconnect", "error", err) + } else { + if u.connections[c.machine] == c { + delete(u.connections, c.machine) + if e, ok := u.endpoints[c.machine]; ok { + seen := c.seen.Load() + e.LastSeenAt = &seen + u.endpoints[c.machine] = e + } + s.broadcast(credential.UserID, u) } - s.broadcast(u) + u.leave() } - u.leave() if c.machine != "" { - _ = s.store.Touch(cleanup, credential.UserID, c.machine, time.UnixMilli(c.seen.Load())) + if err := s.store.Touch(cleanup, credential.UserID, c.machine, time.UnixMilli(c.seen.Load())); err != nil { + s.logger.Warn("端点断开时刷新在线时间失败", "user_id", credential.UserID, "machine_id", c.machine, "operation", "disconnect_touch", "error", err) + } } s.mu.Lock() delete(s.connections, c) @@ -258,6 +290,9 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { kind, data, err := ws.Read(ctx) helloTimer.Stop() if err != nil { + if !c.dead.Load() && !c.expected(err) { + s.logger.Warn("端点初始消息读取失败", "user_id", credential.UserID, "operation", "handshake_read", "error", err) + } s.stats.handshakes.Add(1) c.stop(1002) return @@ -277,7 +312,7 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { if errors.As(err, &f) { name = f.code } - c.write(errorFrame(name, "")) + c.sendError(name, "") s.stats.handshakes.Add(1) c.stop(1002) return @@ -285,10 +320,16 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { checkCtx, checkCancel := context.WithTimeout(ctx, s.timeDB) defer checkCancel() if err = u.enter(checkCtx); err != nil { + if !c.dead.Load() && !c.expected(err) { + s.logger.Warn("端点握手获取锁失败", "user_id", credential.UserID, "operation", "handshake_lock", "error", err) + } c.stop(1013) return } - code, _ := s.check(checkCtx, credential) + code, checkErr := s.check(checkCtx, credential) + if checkErr != nil && !c.dead.Load() && !c.expected(checkErr) { + s.logger.Warn("端点握手鉴权失败", "user_id", credential.UserID, "operation", "handshake_verify", "error", checkErr) + } if code != 0 || s.draining.Load() || c.dead.Load() { u.leave() if code == 0 { @@ -314,16 +355,19 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) { } u.connections[c.machine] = c u.endpoints[c.machine] = e - s.broadcast(u) + s.broadcast(credential.UserID, u) } } u.leave() if err != nil { + if _, name := status(err); name == "service_unavailable" && !c.dead.Load() && !c.expected(err) { + s.logger.Warn("端点注册失败", "user_id", credential.UserID, "machine_id", h.MachineID, "operation", "register", "error", err) + } _, name := status(err) if name == "endpoint_not_found" { name = "unauthorized" } - c.write(errorFrame(name, "")) + c.sendError(name, "") s.stats.handshakes.Add(1) code := 1013 if name == "unauthorized" { diff --git a/monkeyai/backend/internal/endpoint/postgres.go b/monkeyai/backend/internal/endpoint/postgres.go index cc8280824..407e59408 100644 --- a/monkeyai/backend/internal/endpoint/postgres.go +++ b/monkeyai/backend/internal/endpoint/postgres.go @@ -3,6 +3,8 @@ package endpoint import ( "context" "errors" + "fmt" + "log/slog" "time" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/audit" @@ -82,12 +84,20 @@ func (p *Postgres) Get(ctx context.Context, user, machine string) (Endpoint, err row, err := sqlc.New(p.pool).Get(ctx, sqlc.GetParams{UserID: user, MachineID: machine}) return fromRow(row), missing(err) } +func rollback(ctx context.Context, tx pgx.Tx, user, operation string) { + if err := tx.Rollback(ctx); err != nil && ctx.Err() == nil && + !errors.Is(err, pgx.ErrTxClosed) && !errors.Is(err, context.Canceled) && + !errors.Is(err, context.DeadlineExceeded) { + slog.Warn("端点事务回滚失败", "user_id", user, "operation", operation, "error_type", fmt.Sprintf("%T", err)) + } +} + func (p *Postgres) Page(ctx context.Context, user string, page, size int) (Page, error) { tx, err := p.pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}) if err != nil { return Page{}, err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, user, "page") q := sqlc.New(tx) total, err := q.Count(ctx, user) if err != nil { @@ -109,7 +119,7 @@ func (p *Postgres) transaction(ctx context.Context, user string, fn func(*sqlc.Q if err != nil { return Endpoint{}, err } - defer tx.Rollback(ctx) + defer rollback(ctx, tx, user, "transaction") q := sqlc.New(tx) if _, err = q.LockUser(ctx, user); err != nil { return Endpoint{}, missing(err) diff --git a/monkeyai/backend/internal/endpoint/protocol.go b/monkeyai/backend/internal/endpoint/protocol.go index 3a7064863..7ae580ae8 100644 --- a/monkeyai/backend/internal/endpoint/protocol.go +++ b/monkeyai/backend/internal/endpoint/protocol.go @@ -226,7 +226,7 @@ func message(data []byte) (Message, error) { return m, nil } -func errorFrame(code, reply string) []byte { +func errorFrame(code, reply string) ([]byte, error) { descriptions := map[string]string{ "invalid_message": "消息格式无效", "unsupported_protocol": "协议版本不兼容", "unauthorized": "凭据失效", "endpoint_revoked": "端点已停用", "endpoint_limit_exceeded": "端点数量达到上限", "target_unavailable": "目标不可用", @@ -242,6 +242,5 @@ func errorFrame(code, reply string) []byte { Reply string `json:"reply_to,omitempty"` Error wireError `json:"error"` }{"error", reply, e} - data, _ := json.Marshal(frame) - return data + return json.Marshal(frame) } diff --git a/monkeyai/backend/internal/endpoint/protocol_test.go b/monkeyai/backend/internal/endpoint/protocol_test.go index 7ed7be4bf..b0a49ec55 100644 --- a/monkeyai/backend/internal/endpoint/protocol_test.go +++ b/monkeyai/backend/internal/endpoint/protocol_test.go @@ -1,8 +1,18 @@ package endpoint import ( + "bytes" + "context" "encoding/json" + "errors" + "io" + "log/slog" + "net/http" + "net/http/httptest" "strings" + "syscall" + + "github.com/jackc/pgx/v5" "testing" ) @@ -87,3 +97,47 @@ func TestOrigin(t *testing.T) { } } } + +func TestEndpointResponseSerializationFailure(t *testing.T) { + w := httptest.NewRecorder() + respond(w, http.StatusOK, make(chan int)) + if w.Code != http.StatusInternalServerError || strings.Contains(w.Body.String(), "chan") { + t.Fatalf("未正确处理序列化失败: %d %q", w.Code, w.Body.String()) + } +} + +func TestConnectionExpectedErrors(t *testing.T) { + c := &connection{ctx: context.Background()} + c.dead.Store(true) + if c.expected(errors.New("forced close failure")) || !c.expected(io.EOF) || !c.expected(context.Canceled) || !c.expected(syscall.ECONNRESET) { + t.Fatal("主动关闭未掩盖意外失败或误报了正常断开") + } +} + +type rollbackTestTx struct { + pgx.Tx + err error +} + +func (tx rollbackTestTx) Rollback(context.Context) error { return tx.err } + +func TestRollbackLoggingSkipsNormalErrorsAndRedactsDetails(t *testing.T) { + var logs bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + + rollback(context.Background(), rollbackTestTx{err: pgx.ErrTxClosed}, "user-1", "page") + ctx, cancel := context.WithCancel(context.Background()) + cancel() + rollback(ctx, rollbackTestTx{err: errors.New("transaction canceled")}, "user-1", "page") + rollback(context.Background(), rollbackTestTx{err: context.Canceled}, "user-1", "page") + if logs.Len() != 0 { + t.Fatalf("正常事务关闭或取消不应记录警告: %s", logs.String()) + } + rollback(context.Background(), rollbackTestTx{err: errors.New("https://example.com/?token=private-token")}, "user-1", "page") + if !strings.Contains(logs.String(), "error_type=") || !strings.Contains(logs.String(), "user_id=user-1") || + strings.Contains(logs.String(), "private-token") || strings.Contains(logs.String(), "token=") { + t.Fatalf("回滚失败日志缺少安全上下文或泄露敏感信息: %s", logs.String()) + } +} diff --git a/monkeyai/backend/internal/endpoint/service.go b/monkeyai/backend/internal/endpoint/service.go index 19d31c216..7202aa6e4 100644 --- a/monkeyai/backend/internal/endpoint/service.go +++ b/monkeyai/backend/internal/endpoint/service.go @@ -145,16 +145,20 @@ func (s *Service) view(u *userState, e Endpoint) Endpoint { } return e } -func (s *Service) broadcast(u *userState) { +func (s *Service) broadcast(userID string, u *userState) { rows := make([]View, 0, len(u.endpoints)) for _, e := range u.endpoints { rows = append(rows, s.view(u, e).View) } sort.Slice(rows, func(i, j int) bool { return rows[i].MachineID < rows[j].MachineID }) - data, _ := json.Marshal(struct { + data, err := json.Marshal(struct { Type string `json:"type"` Endpoints []View `json:"endpoints"` }{"directory.snapshot", rows}) + if err != nil { + s.logger.Error("序列化端点目录失败", "user_id", userID, "operation", "broadcast", "error", err) + return + } for _, c := range u.connections { c.snapshot(data) } @@ -163,7 +167,10 @@ func (s *Service) route(c *connection, m Message, size int) { u := c.user ctx, cancel := context.WithTimeout(c.ctx, s.timeDB) defer cancel() - if u.enter(ctx) != nil { + if err := u.enter(ctx); err != nil { + if !c.dead.Load() && !c.expected(err) { + s.logger.Warn("端点路由获取锁失败", "user_id", c.credential.UserID, "machine_id", c.machine, "operation", "route", "error", err) + } c.stop(1013) return } @@ -205,6 +212,7 @@ func (s *Service) route(c *connection, m Message, size int) { m.RoutedAt = now.UnixMilli() data, err := json.Marshal(m) if err != nil { + s.logger.Error("序列化端点转发消息失败", "user_id", c.credential.UserID, "machine_id", c.machine, "message_id", m.ID, "operation", "route", "error", err) c.reject("invalid_message", m.ID) return } diff --git a/monkeyai/backend/internal/expert/service.go b/monkeyai/backend/internal/expert/service.go index 026530b92..993f5ce24 100644 --- a/monkeyai/backend/internal/expert/service.go +++ b/monkeyai/backend/internal/expert/service.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "log/slog" "net/http" "strings" @@ -102,9 +103,16 @@ func (s *Service) RegisterAgent(r chi.Router) { } func connectorLinks(v any) []resource.Object { - b, _ := json.Marshal(v) + b, err := json.Marshal(v) + if err != nil { + slog.Error("专家连接依赖编码失败", "error", err) + return nil + } out := []resource.Object{} - _ = json.Unmarshal(b, &out) + if err := json.Unmarshal(b, &out); err != nil { + slog.Error("专家连接依赖解码失败", "error", err) + return nil + } return out } diff --git a/monkeyai/backend/internal/group/move.go b/monkeyai/backend/internal/group/move.go index ee69e682f..53e96d2dc 100644 --- a/monkeyai/backend/internal/group/move.go +++ b/monkeyai/backend/internal/group/move.go @@ -44,7 +44,7 @@ func (s *Service) Move(ctx context.Context, actor string, in MoveInput) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackGroup(ctx, tx, "移动分组或成员", in.TargetID) var target *string if in.TargetID != rootgroup.ID { if _, err := get(ctx, tx, in.TargetID); err != nil { diff --git a/monkeyai/backend/internal/group/service.go b/monkeyai/backend/internal/group/service.go index 206b05b05..edcd53230 100644 --- a/monkeyai/backend/internal/group/service.go +++ b/monkeyai/backend/internal/group/service.go @@ -3,6 +3,8 @@ package group import ( "context" "encoding/json" + "errors" + "log/slog" "slices" "strings" "unicode/utf8" @@ -92,12 +94,12 @@ func (s *Service) begin(ctx context.Context) (pgx.Tx, error) { } // 串行化分组写入,避免并发移动绕过祖先校验形成环。 if _, err = sqlc.New(tx).LockGroups(ctx); err != nil { - _ = tx.Rollback(ctx) + rollbackGroup(ctx, tx, "初始化分组事务", "") return nil, err } if s.accounts != nil { if err = s.accounts.PreserveAccounts(ctx, tx); err != nil { - _ = tx.Rollback(ctx) + rollbackGroup(ctx, tx, "保全分组关联账户", "") return nil, err } } @@ -135,7 +137,7 @@ func (s *Service) Save(ctx context.Context, actor, id string, in Input) (Group, if err != nil { return Group{}, err } - defer tx.Rollback(ctx) + defer rollbackGroup(ctx, tx, "保存分组", id) var group Group if !create { group, err = get(ctx, tx, id) @@ -205,7 +207,7 @@ func (s *Service) SetMembers(ctx context.Context, actor, id string, ids []string if err != nil { return Group{}, err } - defer tx.Rollback(ctx) + defer rollbackGroup(ctx, tx, "更新分组成员", id) group, err := get(ctx, tx, id) if err != nil { return Group{}, err @@ -243,7 +245,7 @@ func (s *Service) Delete(ctx context.Context, actor, id string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackGroup(ctx, tx, "删除分组", id) if _, err := get(ctx, tx, id); err != nil { return err } @@ -270,3 +272,9 @@ func (s *Service) Delete(ctx context.Context, actor, id string) error { } return tx.Commit(ctx) } + +func rollbackGroup(ctx context.Context, tx pgx.Tx, operation, id string) { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚分组事务失败", "operation", operation, "group_id", id, "error", err) + } +} diff --git a/monkeyai/backend/internal/group/service_test.go b/monkeyai/backend/internal/group/service_test.go index 0bf6f9261..d92d62e7c 100644 --- a/monkeyai/backend/internal/group/service_test.go +++ b/monkeyai/backend/internal/group/service_test.go @@ -75,7 +75,10 @@ func TestGroups(t *testing.T) { router.Use(identity.NewService(pool, nil, "http://localhost").RequireAdmin) service.RegisterAdmin(router) call := func(method, path string, body any, token string) *httptest.ResponseRecorder { - data, _ := json.Marshal(body) + data, err := json.Marshal(body) + if err != nil { + t.Fatal(err) + } req := httptest.NewRequest(method, path, bytes.NewReader(data)) if token != "" { req.AddCookie(&http.Cookie{Name: "monkeyai_session", Value: token}) @@ -92,7 +95,9 @@ func TestGroups(t *testing.T) { } var group Group if status != 204 { - _ = json.Unmarshal(response.Body.Bytes(), &group) + if err := json.Unmarshal(response.Body.Bytes(), &group); err != nil { + t.Fatal(err) + } } return group } diff --git a/monkeyai/backend/internal/httpapi/cache.go b/monkeyai/backend/internal/httpapi/cache.go index d810a7e6c..75f281cb4 100644 --- a/monkeyai/backend/internal/httpapi/cache.go +++ b/monkeyai/backend/internal/httpapi/cache.go @@ -4,9 +4,12 @@ import ( "crypto/sha256" "encoding/hex" "encoding/json" + "log/slog" "maps" "net/http" "strings" + + "github.com/go-chi/chi/v5/middleware" ) func CachedJSON(w http.ResponseWriter, r *http.Request, value map[string]any) error { @@ -33,6 +36,8 @@ func CachedJSON(w http.ResponseWriter, r *http.Request, value map[string]any) er } } w.Header().Set("Content-Type", "application/json; charset=utf-8") - _, err = w.Write(encoded) - return err + if _, err := w.Write(encoded); err != nil { + slog.ErrorContext(r.Context(), "缓存响应写入失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + } + return nil } diff --git a/monkeyai/backend/internal/httpapi/cache_test.go b/monkeyai/backend/internal/httpapi/cache_test.go index 694fbede1..405bfe855 100644 --- a/monkeyai/backend/internal/httpapi/cache_test.go +++ b/monkeyai/backend/internal/httpapi/cache_test.go @@ -7,6 +7,13 @@ import ( "testing" ) +func TestCachedJSONWriteFailureDoesNotRequestSecondResponse(t *testing.T) { + w := brokenWriter{httptest.NewRecorder()} + if err := CachedJSON(w, httptest.NewRequest(http.MethodGet, "/rules", nil), map[string]any{"rules": []string{}}); err != nil { + t.Fatalf("响应已开始后不应返回写入错误: %v", err) + } +} + func TestCachedJSON(t *testing.T) { value := map[string]any{"rules": []string{"rule-1"}} read := func(match string) *httptest.ResponseRecorder { diff --git a/monkeyai/backend/internal/httpapi/router.go b/monkeyai/backend/internal/httpapi/router.go index 8641f85ea..c5782801b 100644 --- a/monkeyai/backend/internal/httpapi/router.go +++ b/monkeyai/backend/internal/httpapi/router.go @@ -19,19 +19,25 @@ func New(logger *slog.Logger, database Pinger, admin, agent, auth http.Handler) router.Use(middleware.RequestID) router.Use(middleware.Recoverer) - router.Get("/healthz", func(w http.ResponseWriter, _ *http.Request) { + router.Get("/healthz", func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json; charset=utf-8") - _, _ = io.WriteString(w, "{\"status\":\"ok\"}\n") + if _, err := io.WriteString(w, "{\"status\":\"ok\"}\n"); err != nil { + logger.ErrorContext(r.Context(), "健康检查响应写入失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + } }) router.Get("/readyz", func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json; charset=utf-8") if err := database.Ping(r.Context()); err != nil { - logger.Error("数据库就绪检查失败", "error", err) + logger.ErrorContext(r.Context(), "数据库就绪检查失败", "request_id", middleware.GetReqID(r.Context()), "error", err) w.WriteHeader(http.StatusServiceUnavailable) - _, _ = io.WriteString(w, "{\"status\":\"unavailable\"}\n") + if _, writeErr := io.WriteString(w, "{\"status\":\"unavailable\"}\n"); writeErr != nil { + logger.ErrorContext(r.Context(), "就绪检查响应写入失败", "request_id", middleware.GetReqID(r.Context()), "error", writeErr) + } return } - _, _ = io.WriteString(w, "{\"status\":\"ok\"}\n") + if _, err := io.WriteString(w, "{\"status\":\"ok\"}\n"); err != nil { + logger.ErrorContext(r.Context(), "就绪检查响应写入失败", "request_id", middleware.GetReqID(r.Context()), "error", err) + } }) router.Mount("/api/admin/v1", admin) diff --git a/monkeyai/backend/internal/httpapi/router_test.go b/monkeyai/backend/internal/httpapi/router_test.go index 5665fc90f..9d7e81724 100644 --- a/monkeyai/backend/internal/httpapi/router_test.go +++ b/monkeyai/backend/internal/httpapi/router_test.go @@ -1,12 +1,14 @@ package httpapi import ( + "bytes" "context" "errors" "io" "log/slog" "net/http" "net/http/httptest" + "strings" "testing" ) @@ -53,6 +55,19 @@ func TestReadyWhenDatabaseUnavailable(t *testing.T) { } } +type brokenWriter struct{ http.ResponseWriter } + +func (brokenWriter) Write([]byte) (int, error) { return 0, errors.New("write failed") } + +func TestHealthWriteFailureLogged(t *testing.T) { + var logs bytes.Buffer + handler := New(slog.New(slog.NewTextHandler(&logs, nil)), stubPinger{}, http.NotFoundHandler(), http.NotFoundHandler(), http.NotFoundHandler()) + handler.ServeHTTP(brokenWriter{httptest.NewRecorder()}, httptest.NewRequest(http.MethodGet, "/healthz", nil)) + if !strings.Contains(logs.String(), "健康检查响应写入失败") || !strings.Contains(logs.String(), "write failed") { + t.Fatalf("缺少写入错误日志: %s", logs.String()) + } +} + func TestPprofNotExposed(t *testing.T) { recorder := httptest.NewRecorder() New( diff --git a/monkeyai/backend/internal/identity/admin.go b/monkeyai/backend/internal/identity/admin.go index e418da38a..f4a5bb092 100644 --- a/monkeyai/backend/internal/identity/admin.go +++ b/monkeyai/backend/internal/identity/admin.go @@ -3,6 +3,7 @@ package identity import ( "encoding/json" "errors" + "log/slog" "net/http" "strings" @@ -20,6 +21,7 @@ func (s *Service) RegisterAdmin(router chi.Router) { router.Get("/users", func(w http.ResponseWriter, r *http.Request) { users, err := s.listUsers(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取用户列表失败", "error", err) writeError(w, http.StatusInternalServerError, "server_error", "读取用户失败") return } @@ -60,6 +62,11 @@ func (s *Service) patchUser(w http.ResponseWriter, r *http.Request) { } user, err := s.updateUser(r.Context(), chi.URLParam(r, "userID"), input.Name, input.Role, input.Status, "") if err != nil { + if !errors.Is(err, ErrNotFound) { + slog.ErrorContext(r.Context(), "更新用户失败", "user_id", chi.URLParam(r, "userID"), "error", err) + writeError(w, http.StatusInternalServerError, "server_error", "更新用户失败") + return + } writeError(w, http.StatusNotFound, "user_not_found", "用户不存在") return } @@ -70,21 +77,28 @@ func (s *Service) resetUserPassword(w http.ResponseWriter, r *http.Request) { w.Header().Set("Cache-Control", "no-store") password, err := generatePassword() if err != nil { + slog.ErrorContext(r.Context(), "生成重置密码失败", "user_id", chi.URLParam(r, "userID"), "error", err) writeError(w, http.StatusInternalServerError, "server_error", "生成密码失败") return } hash, err := hashPassword(password) if err != nil { + slog.ErrorContext(r.Context(), "重置用户密码失败", "user_id", chi.URLParam(r, "userID"), "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } ctx := r.Context() tx, err := s.db.Begin(ctx) if err != nil { + slog.ErrorContext(r.Context(), "重置用户密码失败", "user_id", chi.URLParam(r, "userID"), "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚重置用户密码事务失败", "user_id", chi.URLParam(r, "userID"), "error", err) + } + }() q := sqlc.New(tx) user, err := q.GetUser(ctx, chi.URLParam(r, "userID")) if errors.Is(err, pgx.ErrNoRows) { @@ -92,22 +106,27 @@ func (s *Service) resetUserPassword(w http.ResponseWriter, r *http.Request) { return } if err != nil { + slog.ErrorContext(r.Context(), "重置用户密码失败", "user_id", chi.URLParam(r, "userID"), "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } if _, err := q.ResetUserPassword(ctx, sqlc.ResetUserPasswordParams{ID: user.ID, PasswordHash: &hash}); err != nil { + slog.ErrorContext(ctx, "更新用户密码失败", "user_id", user.ID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } if err := revokePasswordAccess(ctx, q, user.ID, user.Email); err != nil { + slog.ErrorContext(ctx, "撤销用户旧凭据失败", "user_id", user.ID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } if err := q.DeleteEmailCode(ctx, sqlc.DeleteEmailCodeParams{Email: user.Email, Purpose: "reset"}); err != nil { + slog.ErrorContext(ctx, "删除用户重置验证码失败", "user_id", user.ID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } if err := tx.Commit(ctx); err != nil { + slog.ErrorContext(ctx, "提交用户密码重置失败", "user_id", user.ID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "重置密码失败") return } diff --git a/monkeyai/backend/internal/identity/admincreate.go b/monkeyai/backend/internal/identity/admincreate.go index cfb943921..5699e89a9 100644 --- a/monkeyai/backend/internal/identity/admincreate.go +++ b/monkeyai/backend/internal/identity/admincreate.go @@ -3,6 +3,8 @@ package identity import ( "context" "errors" + "github.com/jackc/pgx/v5" + "log/slog" "slices" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity/sqlc" @@ -40,7 +42,11 @@ func (s *Service) insertUserWithGroups(ctx context.Context, actor string, input if err != nil { return User{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚创建用户事务失败", "actor_id", actor, "error", err) + } + }() q := sqlc.New(tx) if len(groupIDs) > 0 { // Coordinate with group deletion and membership replacement so validation diff --git a/monkeyai/backend/internal/identity/branding.go b/monkeyai/backend/internal/identity/branding.go index 983ecf093..0213c3b34 100644 --- a/monkeyai/backend/internal/identity/branding.go +++ b/monkeyai/backend/internal/identity/branding.go @@ -2,6 +2,7 @@ package identity import ( "encoding/json" + "log/slog" "net/http" "strings" ) @@ -23,9 +24,13 @@ func (s *Service) branding(w http.ResponseWriter, r *http.Request) { } if s.settings != nil { value, err := s.settings.GetValue(r.Context(), "branding") - if err == nil { + if err != nil { + slog.ErrorContext(r.Context(), "读取品牌设置失败,使用默认值", "error", err) + } else { var configured publicBranding - if json.Unmarshal(value, &configured) == nil { + if err := json.Unmarshal(value, &configured); err != nil { + slog.ErrorContext(r.Context(), "解析品牌设置失败,使用默认值", "error", err) + } else { if workspaceName := strings.TrimSpace(configured.WorkspaceName); workspaceName != "" { branding.WorkspaceName = workspaceName } diff --git a/monkeyai/backend/internal/identity/email.go b/monkeyai/backend/internal/identity/email.go index 518a24087..59fb9d2de 100644 --- a/monkeyai/backend/internal/identity/email.go +++ b/monkeyai/backend/internal/identity/email.go @@ -7,6 +7,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "math/big" "net" "net/http" @@ -42,6 +43,7 @@ func (s *Service) loginMethods(ctx context.Context) (loginMethods, error) { func (s *Service) methods(w http.ResponseWriter, r *http.Request) { methods, err := s.loginMethods(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取邮件认证设置失败", "error", err) writeError(w, 503, "settings_unavailable", "认证配置不可用") return } @@ -51,6 +53,7 @@ func (s *Service) methods(w http.ResponseWriter, r *http.Request) { func (s *Service) allowEmail(w http.ResponseWriter, r *http.Request, purpose string) (loginMethods, bool) { methods, err := s.loginMethods(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取邮件认证设置失败", "error", err) writeError(w, 503, "settings_unavailable", "认证配置不可用") return methods, false } @@ -120,6 +123,7 @@ func (s *Service) sendCode(w http.ResponseWriter, r *http.Request) { return } if err != nil { + slog.ErrorContext(r.Context(), "保留验证码失败", "purpose", input.Purpose, "error", err) writeError(w, 500, "server_error", "发送验证码失败") return } @@ -130,6 +134,7 @@ func (s *Service) sendCode(w http.ResponseWriter, r *http.Request) { eligible = true } if lookupErr != nil && !errors.Is(lookupErr, pgx.ErrNoRows) { + slog.ErrorContext(r.Context(), "查询验证码接收用户失败", "purpose", input.Purpose, "error", lookupErr) writeError(w, 500, "server_error", "发送验证码失败") return } @@ -137,10 +142,12 @@ func (s *Service) sendCode(w http.ResponseWriter, r *http.Request) { labels := map[string]string{"login": "登录", "reset": "重置密码"} err = s.email.Send(r.Context(), input.Email, "MonkeyAI "+labels[input.Purpose]+"验证码", fmt.Sprintf("你的%s验证码为:%s\n\n验证码 10 分钟内有效,仅可使用一次。如非本人操作,请忽略此邮件。", labels[input.Purpose], code)) if err != nil { + slog.ErrorContext(r.Context(), "发送验证码邮件失败", "purpose", input.Purpose, "error", err) writeError(w, 502, "email_failed", "邮件发送失败,请稍后重试或联系管理员") return } if err := sqlc.New(s.db).ReadyEmailCode(r.Context(), sqlc.ReadyEmailCodeParams{Email: input.Email, Purpose: input.Purpose, CodeHash: emailCodeHash(input, code)}); err != nil { + slog.ErrorContext(r.Context(), "激活验证码失败", "purpose", input.Purpose, "error", err) writeError(w, 500, "server_error", "发送验证码失败") return } @@ -160,7 +167,11 @@ func (s *Service) reserveCode(ctx context.Context, input emailInput, code, ipHas if err != nil { return err } - defer tx.Rollback(ctx) + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚邮件验证码事务失败", "purpose", input.Purpose, "error", err) + } + }() q := sqlc.New(tx) if err := q.LockEmailDelivery(ctx); err != nil { return err @@ -250,14 +261,20 @@ func (s *Service) completeEmail(w http.ResponseWriter, r *http.Request, purpose ctx := r.Context() tx, err := s.db.Begin(ctx) if err != nil { + slog.ErrorContext(ctx, "启动验证码校验事务失败", "purpose", purpose, "error", err) writeError(w, 500, "server_error", "认证失败") return } - defer tx.Rollback(ctx) + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚邮件验证码事务失败", "purpose", input.Purpose, "error", err) + } + }() if err := consumeEmailCode(ctx, tx, input); err != nil { if errors.Is(err, errInvalidCode) { writeError(w, 400, "invalid_code", errInvalidCode.Error()) } else { + slog.ErrorContext(ctx, "校验邮件验证码失败", "purpose", purpose, "error", err) writeError(w, 500, "server_error", "认证失败") } return @@ -268,6 +285,7 @@ func (s *Service) completeEmail(w http.ResponseWriter, r *http.Request, purpose case "reset": hash, hashErr := hashPassword(input.Password) if hashErr != nil { + slog.ErrorContext(ctx, "生成重置密码哈希失败", "error", hashErr) writeError(w, 500, "server_error", "密码重置失败") return } @@ -276,6 +294,9 @@ func (s *Service) completeEmail(w http.ResponseWriter, r *http.Request, purpose resetErr = revokePasswordAccess(ctx, q, id, input.Email) } if resetErr != nil { + if !errors.Is(resetErr, pgx.ErrNoRows) { + slog.ErrorContext(ctx, "重置密码并撤销旧凭据失败", "error", resetErr) + } writeError(w, 400, "reset_failed", "密码重置失败,请重新获取验证码") return } @@ -284,18 +305,23 @@ func (s *Service) completeEmail(w http.ResponseWriter, r *http.Request, purpose admin := strings.HasPrefix(r.URL.Path, "/admin/") || strings.Contains(r.URL.Path, "/v1/admin/") if errors.Is(lookupErr, pgx.ErrNoRows) && methods.EmailCodeAutoRegistrationEnabled && !admin { if err := q.CreateEmailUser(ctx, sqlc.CreateEmailUserParams{Name: input.Email, Email: input.Email}); err != nil { + slog.ErrorContext(ctx, "自动注册邮件登录用户失败", "error", err) writeError(w, 500, "server_error", "创建账号失败") return } row, lookupErr = q.GetUserByEmail(ctx, input.Email) } if lookupErr != nil || row.Status != "active" || admin && row.Role != "admin" { + if lookupErr != nil && !errors.Is(lookupErr, pgx.ErrNoRows) { + slog.ErrorContext(ctx, "查询验证码登录用户失败", "error", lookupErr) + } writeError(w, 401, "invalid_credentials", "账号不可用于此登录入口") return } user = User{ID: row.ID, Name: row.Name, Email: row.Email, AvatarURL: row.AvatarUrl, Role: row.Role, Status: row.Status, JoinedAt: row.JoinedAt} } if err := tx.Commit(ctx); err != nil { + slog.ErrorContext(ctx, "提交邮件认证事务失败", "purpose", purpose, "error", err) writeError(w, 500, "server_error", "认证失败") return } diff --git a/monkeyai/backend/internal/identity/email_test.go b/monkeyai/backend/internal/identity/email_test.go index bf5968378..4b655afa1 100644 --- a/monkeyai/backend/internal/identity/email_test.go +++ b/monkeyai/backend/internal/identity/email_test.go @@ -430,7 +430,11 @@ func TestConcurrentEmailCode(t *testing.T) { consumed <- err return } - defer tx.Rollback(t.Context()) + defer func() { + if err := tx.Rollback(t.Context()); err != nil && !errors.Is(err, pgx.ErrTxClosed) { + t.Error(err) + } + }() err = consumeEmailCode(t.Context(), tx, input) if err == nil { err = tx.Commit(t.Context()) diff --git a/monkeyai/backend/internal/identity/middleware.go b/monkeyai/backend/internal/identity/middleware.go index 6b78ab5be..0e6f81143 100644 --- a/monkeyai/backend/internal/identity/middleware.go +++ b/monkeyai/backend/internal/identity/middleware.go @@ -27,6 +27,9 @@ func (s *Service) BrowserUser(r *http.Request) (User, bool) { return User{}, false } user, _, err := s.userByBrowserToken(r.Context(), tokenHash(cookie.Value)) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(r.Context(), "查询浏览器会话失败", "error", err) + } return user, err == nil } @@ -39,6 +42,9 @@ func (s *Service) RequireAdmin(next http.Handler) http.Handler { } user, _, err := s.userByBrowserToken(r.Context(), tokenHash(cookie.Value)) if err != nil { + if !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(r.Context(), "查询管理员会话失败", "error", err) + } writeError(w, http.StatusUnauthorized, "unauthorized", "请先登录") return } @@ -106,7 +112,9 @@ func (s *Service) RequireAgent(next http.Handler) http.Handler { func writeJSON(w http.ResponseWriter, status int, value any) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(value) + if err := json.NewEncoder(w).Encode(value); err != nil { + slog.Error("写入身份服务 HTTP 响应失败", "status", status, "error", err) + } } func writeError(w http.ResponseWriter, status int, code, message string) { diff --git a/monkeyai/backend/internal/identity/oauth.go b/monkeyai/backend/internal/identity/oauth.go index da41cff02..dd8de9d46 100644 --- a/monkeyai/backend/internal/identity/oauth.go +++ b/monkeyai/backend/internal/identity/oauth.go @@ -1,11 +1,14 @@ package identity import ( + "context" "errors" + "log/slog" "net/http" "net/url" "github.com/go-chi/chi/v5" + "github.com/jackc/pgx/v5" ) func (s *Service) OAuthRouter() http.Handler { @@ -98,7 +101,9 @@ func (s *Service) token(w http.ResponseWriter, r *http.Request) { func (s *Service) revoke(w http.ResponseWriter, r *http.Request) { r.Body = http.MaxBytesReader(w, r.Body, 64<<10) if err := r.ParseForm(); err == nil { - _ = s.revokeToken(r.Context(), tokenHash(r.Form.Get("token")), r.Form.Get("client_id")) + if err := s.revokeToken(r.Context(), tokenHash(r.Form.Get("token")), r.Form.Get("client_id")); err != nil { + slog.ErrorContext(r.Context(), "撤销 OAuth 令牌失败", "client_id", r.Form.Get("client_id"), "error", err) + } } w.WriteHeader(http.StatusOK) } @@ -106,6 +111,7 @@ func (s *Service) revoke(w http.ResponseWriter, r *http.Request) { func (s *Service) writeOAuthError(w http.ResponseWriter, err error) { var protocol protocolError if !errors.As(err, &protocol) { + slog.Error("OAuth 请求处理失败", "error", err) protocol = protocolError{Code: "server_error", Description: "服务暂时不可用"} } w.Header().Set("Cache-Control", "no-store") @@ -115,6 +121,7 @@ func (s *Service) writeOAuthError(w http.ResponseWriter, err error) { func (s *Service) providers(w http.ResponseWriter, r *http.Request) { connections, err := s.connections(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取登录方式失败", "error", err) writeError(w, http.StatusServiceUnavailable, "settings_unavailable", "认证配置不可用") return } @@ -144,7 +151,9 @@ func nullableUser(user User, ok bool) any { func (s *Service) logout(w http.ResponseWriter, r *http.Request) { if cookie, err := r.Cookie(sessionCookie); err == nil { - _ = s.revokeBrowserSession(r.Context(), tokenHash(cookie.Value)) + if err := s.revokeBrowserSession(r.Context(), tokenHash(cookie.Value)); err != nil { + slog.ErrorContext(r.Context(), "撤销浏览器会话失败", "error", err) + } } http.SetCookie(w, &http.Cookie{Name: sessionCookie, Value: "", Path: "/", MaxAge: -1, HttpOnly: true, Secure: s.secureCookie, SameSite: http.SameSiteLaxMode}) w.WriteHeader(http.StatusNoContent) @@ -153,10 +162,18 @@ func (s *Service) logout(w http.ResponseWriter, r *http.Request) { func (s *Service) clientRequest(w http.ResponseWriter, r *http.Request) { request, err := s.authorizationRequest(r.Context(), chi.URLParam(r, "requestID")) if err != nil || request.CompletedAt != nil || !s.now().Before(request.ExpiresAt) { + if err != nil && !errors.Is(err, ErrNotFound) { + slog.ErrorContext(r.Context(), "查询客户端授权请求失败", "request_id", chi.URLParam(r, "requestID"), "error", err) + } writeError(w, http.StatusNotFound, "request_not_found", "授权请求不存在、已完成或已过期") return } - connections, _ := s.connections(r.Context()) + connections, err := s.connections(r.Context()) + if err != nil { + slog.ErrorContext(r.Context(), "读取客户端授权登录方式失败", "request_id", request.ID, "error", err) + writeError(w, http.StatusServiceUnavailable, "settings_unavailable", "认证配置不可用") + return + } user, authenticated := s.BrowserUser(r) client := Clients[request.ClientID] writeJSON(w, http.StatusOK, map[string]any{ @@ -194,6 +211,9 @@ func (s *Service) startUpstream(w http.ResponseWriter, r *http.Request) { } request, err := s.authorizationRequest(r.Context(), requestID) if err != nil || request.CompletedAt != nil || !s.now().Before(request.ExpiresAt) { + if err != nil && !errors.Is(err, ErrNotFound) { + slog.ErrorContext(r.Context(), "查询上游登录授权请求失败", "request_id", requestID, "error", err) + } writeError(w, http.StatusBadRequest, "request_unavailable", "授权请求无效") return } @@ -208,21 +228,27 @@ func (s *Service) beginUpstream(w http.ResponseWriter, r *http.Request, requestI connectionID := chi.URLParam(r, "connectionID") connection, err := s.connection(r.Context(), connectionID) if err != nil { + if !errors.Is(err, ErrNotFound) { + slog.ErrorContext(r.Context(), "读取上游连接失败", "connection_id", connectionID, "error", err) + } writeError(w, http.StatusNotFound, "provider_not_found", "登录方式不存在") return } state, err := randomToken(32) if err != nil { + slog.ErrorContext(r.Context(), "准备上游登录失败", "connection_id", connectionID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "无法发起登录") return } if err := s.createLoginState(r.Context(), tokenHash(state), connectionID, requestID, purpose, s.now().Add(s.requestTTL)); err != nil { + slog.ErrorContext(r.Context(), "创建上游登录状态失败", "connection_id", connectionID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "无法发起登录") return } target, err := s.upstreamAuthorizeURL(r.Context(), connection, state) if err != nil { - writeError(w, http.StatusBadGateway, "provider_unavailable", err.Error()) + logUpstreamFailure(r.Context(), "构造上游授权地址", connectionID, err) + writeError(w, http.StatusBadGateway, "provider_unavailable", "登录方式暂不可用") return } http.Redirect(w, r, target, http.StatusFound) @@ -231,6 +257,9 @@ func (s *Service) beginUpstream(w http.ResponseWriter, r *http.Request, requestI func (s *Service) upstreamCallback(w http.ResponseWriter, r *http.Request) { state, err := s.consumeLoginState(r.Context(), tokenHash(r.URL.Query().Get("state"))) if err != nil { + if !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(r.Context(), "读取上游登录状态失败", "error", err) + } http.Redirect(w, r, s.clientLoginURL("", "oauth_callback"), http.StatusFound) return } @@ -240,11 +269,15 @@ func (s *Service) upstreamCallback(w http.ResponseWriter, r *http.Request) { } connection, err := s.connection(r.Context(), state.ConnectionID) if err != nil { + if !errors.Is(err, ErrNotFound) { + slog.ErrorContext(r.Context(), "读取上游登录配置失败", "connection_id", state.ConnectionID, "error", err) + } http.Redirect(w, r, s.upstreamResultURL(state, "provider_unavailable"), http.StatusFound) return } profile, err := s.exchangeUpstream(r.Context(), connection, r.URL.Query().Get("code")) if err != nil { + logUpstreamFailure(r.Context(), "交换上游身份", state.ConnectionID, err) http.Redirect(w, r, s.upstreamResultURL(state, "oauth_exchange"), http.StatusFound) return } @@ -262,7 +295,11 @@ func (s *Service) upstreamCallback(w http.ResponseWriter, r *http.Request) { return } token, err := randomToken(32) - if err != nil || s.createBrowserSession(r.Context(), user.ID, tokenHash(token), "oauth", s.now().Add(s.sessionTTL)) != nil { + if err == nil { + err = s.createBrowserSession(r.Context(), user.ID, tokenHash(token), "oauth", s.now().Add(s.sessionTTL)) + } + if err != nil { + slog.ErrorContext(r.Context(), "创建上游登录会话失败", "user_id", user.ID, "error", err) http.Redirect(w, r, s.upstreamResultURL(state, "session_failed"), http.StatusFound) return } @@ -295,3 +332,26 @@ func (s *Service) clientLoginURL(requestID, errorCode string) string { } return s.publicURL + "/client-login?" + query.Encode() } + +func logUpstreamFailure(ctx context.Context, operation, connectionID string, err error) { + var response *upstreamHTTPError + if errors.As(err, &response) { + slog.ErrorContext(ctx, "上游 OAuth 操作失败", "operation", operation, "connection_id", connectionID, "reason", "http_status", "status", response.status) + return + } + reason := "invalid_upstream_response" + var urlErr *url.Error + switch { + case errors.Is(err, context.Canceled): + reason = "request_canceled" + case errors.Is(err, context.DeadlineExceeded): + reason = "request_timeout" + case errors.As(err, &urlErr): + if urlErr.Op == "parse" { + reason = "invalid_upstream_url" + } else { + reason = "request_failed" + } + } + slog.ErrorContext(ctx, "上游 OAuth 操作失败", "operation", operation, "connection_id", connectionID, "reason", reason) +} diff --git a/monkeyai/backend/internal/identity/password.go b/monkeyai/backend/internal/identity/password.go index b1973f310..aff06468f 100644 --- a/monkeyai/backend/internal/identity/password.go +++ b/monkeyai/backend/internal/identity/password.go @@ -10,6 +10,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "math/big" "net/http" "net/mail" @@ -18,8 +19,11 @@ import ( "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/audit" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity/sqlc" + "github.com/jackc/pgx/v5" ) +var errPasswordNotSet = errors.New("未设置密码") + const ( passwordIterations = 600_000 dummyPasswordHash = "$pbkdf2-sha256$600000$AAAAAAAAAAAAAAAAAAAAAA$AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA" @@ -53,7 +57,11 @@ func (s *Service) EnsureInitialAdmin(ctx context.Context, name, email, password if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚初始化管理员事务失败", "error", err) + } + }() if _, err := sqlc.New(tx).LockInitialAdmin(ctx); err != nil { return err } @@ -86,6 +94,7 @@ func (s *Service) EnsureInitialAdmin(ctx context.Context, name, email, password func (s *Service) passwordLogin(w http.ResponseWriter, r *http.Request) { methods, err := s.loginMethods(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取密码认证设置失败", "error", err) writeError(w, 503, "settings_unavailable", "认证配置不可用") return } @@ -107,13 +116,16 @@ func (s *Service) passwordLogin(w http.ResponseWriter, r *http.Request) { var passwordHash string row, err := sqlc.New(s.db).GetPasswordUser(r.Context(), sqlc.GetPasswordUserParams{Email: input.Email, AdminOnly: adminOnly}) if err == nil && row.PasswordHash == nil { - err = errors.New("未设置密码") + err = errPasswordNotSet } if err == nil { user = User{ID: row.ID, Name: row.Name, Email: row.Email, AvatarURL: row.AvatarUrl, Role: row.Role, Status: row.Status, JoinedAt: row.JoinedAt, LastLoginAt: row.LastLoginAt} passwordHash = *row.PasswordHash } if err != nil { + if !errors.Is(err, pgx.ErrNoRows) && !errors.Is(err, errPasswordNotSet) { + slog.ErrorContext(r.Context(), "查询密码登录用户失败", "error", err) + } passwordHash = dummyPasswordHash } valid := verifyPassword(input.Password, passwordHash) @@ -126,11 +138,18 @@ func (s *Service) passwordLogin(w http.ResponseWriter, r *http.Request) { func (s *Service) loginSession(w http.ResponseWriter, r *http.Request, user User, method string) { if _, err := sqlc.New(s.db).TouchLogin(r.Context(), user.ID); err != nil { + slog.ErrorContext(r.Context(), "更新登录时间失败", "user_id", user.ID, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "登录失败") return } token, err := randomToken(32) - if err != nil || s.createBrowserSession(r.Context(), user.ID, tokenHash(token), method, s.now().Add(s.sessionTTL)) != nil { + if err != nil { + slog.ErrorContext(r.Context(), "生成登录会话令牌失败", "user_id", user.ID, "error", err) + writeError(w, http.StatusInternalServerError, "server_error", "登录失败") + return + } + if err := s.createBrowserSession(r.Context(), user.ID, tokenHash(token), method, s.now().Add(s.sessionTTL)); err != nil { + slog.ErrorContext(r.Context(), "创建登录会话失败", "user_id", user.ID, "method", method, "error", err) writeError(w, http.StatusInternalServerError, "server_error", "登录失败") return } diff --git a/monkeyai/backend/internal/identity/postgres.go b/monkeyai/backend/internal/identity/postgres.go index e1586bbdd..3197c8049 100644 --- a/monkeyai/backend/internal/identity/postgres.go +++ b/monkeyai/backend/internal/identity/postgres.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "log/slog" "time" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity/sqlc" @@ -45,10 +46,17 @@ func (s *Service) storeAuthorizationCode(ctx context.Context, request Authorizat if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚授权请求事务失败", "request_id", request.ID, "error", err) + } + }() result, err := sqlc.New(tx).CompleteAuthorizationRequest(ctx, request.ID) - if err != nil || result.RowsAffected() != 1 { + if err != nil { + return fmt.Errorf("完成授权请求 %s: %w", request.ID, err) + } + if result.RowsAffected() != 1 { return errors.New("授权请求已完成或已过期") } _, err = sqlc.New(tx).CreateAuthorizationCode(ctx, sqlc.CreateAuthorizationCodeParams{ @@ -76,16 +84,25 @@ func (s *Service) authorizationCode(ctx context.Context, hash string) (Authoriza return code, err } +var errAuthorizationCodeUsed = errors.New("授权码已被使用") + func (s *Service) redeemCodeAndStoreToken(ctx context.Context, codeID, userID, clientID, accessHash, refreshHash string, accessExpiry, refreshExpiry time.Time) error { tx, err := s.db.Begin(ctx) if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚授权码事务失败", "code_id", codeID, "error", err) + } + }() result, err := sqlc.New(tx).RedeemAuthorizationCode(ctx, codeID) - if err != nil || result.RowsAffected() != 1 { - return errors.New("授权码已被使用") + if err != nil { + return fmt.Errorf("兑换授权码 %s: %w", codeID, err) + } + if result.RowsAffected() != 1 { + return errAuthorizationCodeUsed } _, err = sqlc.New(tx).CreateToken(ctx, sqlc.CreateTokenParams{ UserID: userID, @@ -106,7 +123,11 @@ func (s *Service) rotateToken(ctx context.Context, oldRefreshHash, clientID, acc if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚刷新令牌事务失败", "client_id", clientID, "error", err) + } + }() var userID string userID, err = sqlc.New(tx).RevokeRefreshToken(ctx, sqlc.RevokeRefreshTokenParams{RefreshTokenHash: oldRefreshHash, ClientID: clientID}) @@ -214,7 +235,11 @@ func (s *Service) updateUser(ctx context.Context, id, name, role, status, passwo if err != nil { return User{}, err } - defer tx.Rollback(ctx) + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚更新用户事务失败", "user_id", id, "error", err) + } + }() if s.accounts != nil { if err = s.accounts.PreserveAccounts(ctx, tx); err != nil { return User{}, err @@ -240,7 +265,11 @@ func (s *Service) upsertIdentity(ctx context.Context, profile upstreamProfile, a if err != nil { return User{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚上游身份事务失败", "provider", profile.Provider, "error", err) + } + }() identityRow, err := sqlc.New(tx).GetIdentityUser(ctx, sqlc.GetIdentityUserParams{Provider: profile.Provider, Issuer: profile.Issuer, ProviderSubject: profile.Subject}) user := User{ID: identityRow.ID, Name: identityRow.Name, Email: identityRow.Email, AvatarURL: identityRow.AvatarUrl, Role: identityRow.Role, Status: identityRow.Status, JoinedAt: identityRow.JoinedAt, LastLoginAt: identityRow.LastLoginAt} diff --git a/monkeyai/backend/internal/identity/service.go b/monkeyai/backend/internal/identity/service.go index a1e930e24..db017cab9 100644 --- a/monkeyai/backend/internal/identity/service.go +++ b/monkeyai/backend/internal/identity/service.go @@ -10,6 +10,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "net/http" "net/url" "regexp" @@ -193,7 +194,10 @@ func (s *Service) CompleteAuthorization(ctx context.Context, requestID, userID s if err := s.storeAuthorizationCode(ctx, request, userID, tokenHash(code), s.now().Add(s.codeTTL)); err != nil { return "", err } - callback, _ := url.Parse(request.RedirectURI) + callback, err := url.Parse(request.RedirectURI) + if err != nil { + return "", fmt.Errorf("解析授权请求 %s 的回调地址: %w", request.ID, err) + } query := callback.Query() query.Set("code", code) query.Set("state", request.State) @@ -209,6 +213,9 @@ func (s *Service) ExchangeCode(ctx context.Context, clientID, redirectURI, code, return Token{}, oauthError("invalid_grant", "code_verifier 无效") } stored, err := s.authorizationCode(ctx, tokenHash(code)) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(ctx, "查询授权码失败", "client_id", clientID, "error", err) + } if err != nil || stored.ClientID != clientID || stored.RedirectURI != redirectURI || stored.RedeemedAt != nil || !s.now().Before(stored.ExpiresAt) { return Token{}, oauthError("invalid_grant", "授权码无效或已过期") } @@ -231,6 +238,9 @@ func (s *Service) ExchangeCode(ctx context.Context, clientID, redirectURI, code, ExpiresIn: int64(s.accessTTL.Seconds()), ExpiresAt: s.now().Add(s.accessTTL), } if err := s.redeemCodeAndStoreToken(ctx, stored.ID, stored.UserID, clientID, tokenHash(access), tokenHash(refresh), token.ExpiresAt, s.now().Add(s.refreshTTL)); err != nil { + if !errors.Is(err, errAuthorizationCodeUsed) { + slog.ErrorContext(ctx, "兑换授权码失败", "code_id", stored.ID, "error", err) + } return Token{}, oauthError("invalid_grant", "授权码已被使用") } return token, nil @@ -250,6 +260,9 @@ func (s *Service) Refresh(ctx context.Context, clientID, refreshToken string) (T } token := Token{AccessToken: access, RefreshToken: refresh, TokenType: "Bearer", ExpiresIn: int64(s.accessTTL.Seconds()), ExpiresAt: s.now().Add(s.accessTTL)} if err := s.rotateToken(ctx, tokenHash(refreshToken), clientID, tokenHash(access), tokenHash(refresh), token.ExpiresAt, s.now().Add(s.refreshTTL)); err != nil { + if !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(ctx, "刷新 OAuth 令牌失败", "client_id", clientID, "error", err) + } return Token{}, oauthError("invalid_grant", "refresh_token 无效或已过期") } return token, nil diff --git a/monkeyai/backend/internal/identity/upstream.go b/monkeyai/backend/internal/identity/upstream.go index c5d007c25..021f32dc2 100644 --- a/monkeyai/backend/internal/identity/upstream.go +++ b/monkeyai/backend/internal/identity/upstream.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "io" + "log/slog" "net/http" "net/url" "strconv" @@ -35,6 +36,15 @@ type authenticationSettings struct { OAuthConnections []OAuthConnection `json:"oauth_connections"` } +type upstreamHTTPError struct { + operation string + status int +} + +func (e *upstreamHTTPError) Error() string { + return fmt.Sprintf("%s: HTTP %d", e.operation, e.status) +} + type providerMetadata struct { AuthorizationEndpoint string `json:"authorization_endpoint"` TokenEndpoint string `json:"token_endpoint"` @@ -113,9 +123,9 @@ func (s *Service) providerURLs(ctx context.Context, connection OAuthConnection) if err != nil { return providerMetadata{}, fmt.Errorf("读取 OIDC 元数据: %w", err) } - defer response.Body.Close() + defer closeUpstreamBody(ctx, response.Body, "读取 OIDC 元数据") if response.StatusCode != http.StatusOK { - return providerMetadata{}, fmt.Errorf("读取 OIDC 元数据: HTTP %d", response.StatusCode) + return providerMetadata{}, &upstreamHTTPError{operation: "读取 OIDC 元数据", status: response.StatusCode} } if err := json.NewDecoder(io.LimitReader(response.Body, 1<<20)).Decode(&metadata); err != nil { return providerMetadata{}, fmt.Errorf("解析 OIDC 元数据: %w", err) @@ -177,10 +187,9 @@ func (s *Service) exchangeUpstream(ctx context.Context, connection OAuthConnecti if err != nil { return upstreamProfile{}, fmt.Errorf("交换上游令牌: %w", err) } - defer response.Body.Close() + defer closeUpstreamBody(ctx, response.Body, "交换上游令牌") if response.StatusCode < 200 || response.StatusCode >= 300 { - body, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - return upstreamProfile{}, fmt.Errorf("交换上游令牌: HTTP %d: %s", response.StatusCode, strings.TrimSpace(string(body))) + return upstreamProfile{}, &upstreamHTTPError{operation: "交换上游令牌", status: response.StatusCode} } var tokenResponse struct { AccessToken string `json:"access_token"` @@ -199,9 +208,9 @@ func (s *Service) exchangeUpstream(ctx context.Context, connection OAuthConnecti if err != nil { return upstreamProfile{}, fmt.Errorf("读取上游用户: %w", err) } - defer response.Body.Close() + defer closeUpstreamBody(ctx, response.Body, "读取上游用户") if response.StatusCode != http.StatusOK { - return upstreamProfile{}, fmt.Errorf("读取上游用户: HTTP %d", response.StatusCode) + return upstreamProfile{}, &upstreamHTTPError{operation: "读取上游用户", status: response.StatusCode} } decoder := json.NewDecoder(io.LimitReader(response.Body, 1<<20)) var profile upstreamProfile @@ -280,3 +289,9 @@ func stringValue(values map[string]any, keys ...string) string { } return "" } + +func closeUpstreamBody(ctx context.Context, body io.ReadCloser, operation string) { + if err := body.Close(); err != nil && ctx.Err() == nil { + slog.ErrorContext(ctx, "关闭上游响应失败", "operation", operation, "error_type", fmt.Sprintf("%T", err)) + } +} diff --git a/monkeyai/backend/internal/identity/upstream_test.go b/monkeyai/backend/internal/identity/upstream_test.go index 9e3c42aae..92bf481f1 100644 --- a/monkeyai/backend/internal/identity/upstream_test.go +++ b/monkeyai/backend/internal/identity/upstream_test.go @@ -1,7 +1,12 @@ package identity import ( + "bytes" + "context" "errors" + "fmt" + "io" + "log/slog" "net/http" "net/http/httptest" "net/url" @@ -26,7 +31,9 @@ func TestBaizhiyunOIDC(t *testing.T) { t.Error("缺少上游令牌") } w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(userinfo)) + if _, err := w.Write([]byte(userinfo)); err != nil { + t.Error(err) + } default: http.NotFound(w, r) } @@ -254,3 +261,125 @@ func TestBaizhiyunEmailBinding(t *testing.T) { }) } } + +func TestExchangeUpstreamDoesNotExposeTokenResponse(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + if _, err := w.Write([]byte(`{"access_token":"private-token","password":"private-password"}`)); err != nil { + t.Error(err) + } + })) + defer upstream.Close() + s := NewService(nil, nil, "https://monkeyai.example") + connection := OAuthConnection{Provider: "oidc", AuthorizationURL: upstream.URL, TokenURL: upstream.URL, UserInfoURL: upstream.URL} + _, err := s.exchangeUpstream(t.Context(), connection, "code") + if err == nil || !strings.Contains(err.Error(), "HTTP 401") || strings.Contains(err.Error(), "private-") { + t.Fatalf("上游敏感响应不得出现在错误中: %v", err) + } +} + +type upstreamTransport func(*http.Request) (*http.Response, error) + +func (f upstreamTransport) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func TestUpstreamFailureLogOmitsCredentialURLs(t *testing.T) { + var output bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&output, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + + s := NewService(nil, nil, "https://monkeyai.example") + connection := OAuthConnection{ + Provider: "oidc", AuthorizationURL: "https://example.com/authorize%zz?client_secret=private-authorize", + TokenURL: "https://example.com/token?client_secret=private-token", UserInfoURL: "https://example.com/userinfo?access_token=private-userinfo", + } + _, err := s.upstreamAuthorizeURL(t.Context(), connection, "state") + if err == nil || !strings.Contains(err.Error(), "private-authorize") { + t.Fatalf("测试应覆盖包含凭据 URL 的解析错误: %v", err) + } + logUpstreamFailure(t.Context(), "构造上游授权地址", "connection-1", err) + if !strings.Contains(output.String(), "invalid_upstream_url") || strings.Contains(output.String(), "private-") { + t.Fatalf("授权地址错误日志泄露凭据: %s", output.String()) + } + + for _, stage := range []string{"token", "userinfo"} { + t.Run(stage, func(t *testing.T) { + output.Reset() + s.client = &http.Client{Transport: upstreamTransport(func(r *http.Request) (*http.Response, error) { + if stage == "userinfo" && r.URL.Path == "/token" { + return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"access_token":"private-access"}`)), Header: make(http.Header)}, nil + } + return nil, fmt.Errorf("transport failure: %s %s", r.URL.String(), r.Header.Get("Authorization")) + })} + _, err := s.exchangeUpstream(t.Context(), connection, "private-code") + if err == nil || !strings.Contains(err.Error(), "private-") { + t.Fatalf("测试应覆盖包含凭据的外部传输错误: %v", err) + } + logUpstreamFailure(t.Context(), "交换上游身份", "connection-1", err) + if !strings.Contains(output.String(), "reason=request_failed") || strings.Contains(output.String(), "private-") { + t.Fatalf("上游传输错误日志泄露凭据: %s", output.String()) + } + }) + } + + output.Reset() + logUpstreamFailure(t.Context(), "交换上游身份", "connection-1", &upstreamHTTPError{operation: "交换上游令牌", status: http.StatusUnauthorized}) + if !strings.Contains(output.String(), "status=401") || strings.Contains(output.String(), "private-") { + t.Fatalf("上游状态码日志不安全: %s", output.String()) + } +} + +type failingCloseBody struct { + io.Reader + err error +} + +func (body failingCloseBody) Close() error { return body.err } + +func TestUpstreamResponseCloseLogsOnlySafeContext(t *testing.T) { + var output bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&output, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + + s := NewService(nil, nil, "https://monkeyai.example") + s.client = &http.Client{Transport: upstreamTransport(func(r *http.Request) (*http.Response, error) { + var content string + switch r.URL.Path { + case "/.well-known/openid-configuration": + content = `{"authorization_endpoint":"https://example.com/authorize","token_endpoint":"https://example.com/token","userinfo_endpoint":"https://example.com/userinfo"}` + case "/token": + content = `{"access_token":"private-access-token"}` + case "/userinfo": + content = `{"sub":"user-1"}` + default: + t.Errorf("意外的上游请求: %s", r.URL.Path) + } + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: failingCloseBody{ + Reader: strings.NewReader(content), err: fmt.Errorf("关闭 %s?access_token=private-close-token 失败", r.URL.Path), + }}, nil + })} + if _, err := s.providerURLs(t.Context(), OAuthConnection{Provider: "oidc", IssuerURL: "https://example.com"}); err != nil { + t.Fatal(err) + } + connection := OAuthConnection{Provider: "oidc", AuthorizationURL: "https://example.com/authorize", TokenURL: "https://example.com/token", UserInfoURL: "https://example.com/userinfo"} + if _, err := s.exchangeUpstream(t.Context(), connection, "code"); err != nil { + t.Fatal(err) + } + for _, operation := range []string{"读取 OIDC 元数据", "交换上游令牌", "读取上游用户"} { + if !strings.Contains(output.String(), operation) { + t.Errorf("缺少 %s 的关闭失败日志: %s", operation, output.String()) + } + } + if strings.Count(output.String(), "关闭上游响应失败") != 3 || strings.Contains(output.String(), "private-") { + t.Fatalf("关闭失败日志缺失或泄露凭据: %s", output.String()) + } + + output.Reset() + ctx, cancel := context.WithCancel(t.Context()) + cancel() + closeUpstreamBody(ctx, failingCloseBody{Reader: strings.NewReader(""), err: errors.New("private-canceled")}, "读取上游用户") + if output.Len() != 0 { + t.Fatalf("取消上下文不应记录关闭错误: %s", output.String()) + } +} diff --git a/monkeyai/backend/internal/imagegen/inputs.go b/monkeyai/backend/internal/imagegen/inputs.go index b50f60b47..4b0156e72 100644 --- a/monkeyai/backend/internal/imagegen/inputs.go +++ b/monkeyai/backend/internal/imagegen/inputs.go @@ -11,6 +11,7 @@ import ( _ "image/jpeg" _ "image/png" "io" + "log/slog" "time" _ "golang.org/x/image/webp" @@ -76,7 +77,9 @@ func (s *Inputs) Upload(ctx context.Context, userID string, data []byte) (Input, if err != nil { cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) defer cancel() - _ = s.storage.Delete(cleanupCtx, key) + if cleanupErr := s.storage.Delete(cleanupCtx, key); cleanupErr != nil { + slog.Error("清理未登记参考图失败", "operation", "delete_input", "file_id", id, "error_type", providerErrorType(cleanupErr)) + } return Input{}, err } return Input{ID: id, MIMEType: mime, Width: config.Width, Height: config.Height, ExpiresAt: expiry}, nil @@ -94,7 +97,11 @@ func (s *Inputs) Read(ctx context.Context, userID, fileID string) (Input, error) if err != nil { return Input{}, err } - defer body.Close() + defer func() { + if closeErr := body.Close(); closeErr != nil && ctx.Err() == nil { + slog.Warn("关闭参考图存储流失败", "operation", "close_reader", "file_id", fileID, "error_type", providerErrorType(closeErr)) + } + }() data, err := io.ReadAll(io.LimitReader(body, maxInputBytes+1)) if err != nil { return Input{}, err diff --git a/monkeyai/backend/internal/imagegen/outputs.go b/monkeyai/backend/internal/imagegen/outputs.go index d1b9aefeb..e3c20da62 100644 --- a/monkeyai/backend/internal/imagegen/outputs.go +++ b/monkeyai/backend/internal/imagegen/outputs.go @@ -9,6 +9,7 @@ import ( "fmt" "image" "io" + "log/slog" "net/http" "time" @@ -91,7 +92,9 @@ func (s *Outputs) Save(ctx context.Context, jobID string, ordinal int32, img Ima if err != nil { cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) defer cancel() - _ = s.storage.Delete(cleanupCtx, key) + if cleanupErr := s.storage.Delete(cleanupCtx, key); cleanupErr != nil { + slog.Error("清理未登记生图结果失败", "operation", "delete_output", "job_id", jobID, "ordinal", ordinal, "error_type", providerErrorType(cleanupErr)) + } return SavedOutput{}, err } if count == 0 { @@ -102,7 +105,9 @@ func (s *Outputs) Save(ctx context.Context, jobID string, ordinal int32, img Ima if existing.Sha256 != hash { cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) defer cancel() - _ = s.storage.Delete(cleanupCtx, key) + if cleanupErr := s.storage.Delete(cleanupCtx, key); cleanupErr != nil { + slog.Error("清理冲突生图结果失败", "operation", "delete_output", "job_id", jobID, "ordinal", ordinal, "error_type", providerErrorType(cleanupErr)) + } return SavedOutput{}, errors.New("生成图片结果不一致") } id = existing.ID @@ -138,7 +143,11 @@ func (s *Outputs) Open(ctx context.Context, userID, outputID string) ([]byte, st if err != nil { return nil, "", err } - defer body.Close() + defer func() { + if closeErr := body.Close(); closeErr != nil && ctx.Err() == nil { + slog.Warn("关闭生图结果存储流失败", "operation", "close_reader", "file_id", outputID, "error_type", providerErrorType(closeErr)) + } + }() data, err := io.ReadAll(io.LimitReader(body, maxOutputBytes+1)) if err != nil { return nil, "", err diff --git a/monkeyai/backend/internal/imagegen/postgres.go b/monkeyai/backend/internal/imagegen/postgres.go index 715f4c9a9..d624c9cf1 100644 --- a/monkeyai/backend/internal/imagegen/postgres.go +++ b/monkeyai/backend/internal/imagegen/postgres.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "time" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/imagegen/sqlc" @@ -91,7 +92,13 @@ func (p *Postgres) LinkInputs(ctx context.Context, jobID string, fileIDs []strin if err != nil { return err } - defer tx.Rollback(ctx) + defer func() { + if rollbackErr := tx.Rollback(ctx); rollbackErr != nil && ctx.Err() == nil && + !errors.Is(rollbackErr, pgx.ErrTxClosed) && !errors.Is(rollbackErr, context.Canceled) && + !errors.Is(rollbackErr, context.DeadlineExceeded) { + slog.Warn("生图参考图关联回滚失败", "job_id", jobID, "operation", "rollback_link_inputs", "error_type", fmt.Sprintf("%T", rollbackErr)) + } + }() for _, fileID := range fileIDs { if _, err := sqlc.New(tx).LinkJobInput(ctx, sqlc.LinkJobInputParams{JobID: jobID, InputID: fileID}); err != nil { if errors.Is(err, pgx.ErrNoRows) { diff --git a/monkeyai/backend/internal/imagegen/service.go b/monkeyai/backend/internal/imagegen/service.go index 97c4b4bfd..aeedbcf84 100644 --- a/monkeyai/backend/internal/imagegen/service.go +++ b/monkeyai/backend/internal/imagegen/service.go @@ -5,7 +5,10 @@ import ( "crypto/sha256" "encoding/hex" "encoding/json" + "errors" + "fmt" "log/slog" + "net/url" "slices" "strings" "sync" @@ -249,14 +252,23 @@ func (s *Service) submit(ctx context.Context, target proxy.Target, operation, re fileIDs = append(fileIDs, mask.FileID) maskInput = &image } - body, _ := json.Marshal(struct { + body, err := json.Marshal(struct { Model, Operation, Prompt, Quality, Aspect string Count uint32 Inputs []string }{requestedModel, operation, prompt, quality, aspect, imageCount, digests}) + if err != nil { + return imageproxy.Task{}, err + } hash := sha256.Sum256(body) - config, _ := json.Marshal(map[string]any{"file_ids": fileIDs}) - pricing, _ := json.Marshal(map[string]string{"unit": unit.String()}) + config, err := json.Marshal(map[string]any{"file_ids": fileIDs}) + if err != nil { + return imageproxy.Task{}, err + } + pricing, err := json.Marshal(map[string]string{"unit": unit.String()}) + if err != nil { + return imageproxy.Task{}, err + } job, created, err := s.jobs.Create(ctx, Job{ UserID: target.UserID, ModelID: item.ID, Provider: string(item.Provider), Operation: operation, RequestHash: hex.EncodeToString(hash[:]), IdempotencyKey: idempotency, RequestedImages: int32(imageCount), @@ -333,14 +345,28 @@ func (s *Service) failUnsubmitted(ctx context.Context, job Job, code string) { } } +func providerErrorType(err error) string { + var urlErr *url.Error + if errors.As(err, &urlErr) { + return "url_error" + } + return fmt.Sprintf("%T", err) +} + func (s *Service) handleResult(ctx context.Context, job Job, result ProviderResult, err error) { if err != nil { + if ctx.Err() == nil { + slog.Warn("生图上游请求失败,任务状态待确认", "job", job.ID, "operation", job.Operation, "error_type", providerErrorType(err)) + } s.markUnknown(ctx, job, "provider_status_unknown") return } switch result.Status { case "pending", "running": - if result.ID == "" || s.jobs.Running(ctx, job.ID, result.ID) != nil { + if result.ID == "" { + s.markUnknown(ctx, job, "provider_job_id_missing") + } else if err := s.jobs.Running(ctx, job.ID, result.ID); err != nil { + slog.Error("记录生图任务运行状态失败", "job", job.ID, "operation", "running", "error", err) s.markUnknown(ctx, job, "provider_job_id_missing") } case "failed": @@ -422,6 +448,8 @@ func (s *Service) Get(ctx context.Context, userID, jobID string) (imageproxy.Tas if job.BillingTransactionID != nil && (status == "succeeded" || status == "failed" || status == "expired") { if amount, err := s.billing.ImageCharge(ctx, *job.BillingTransactionID, userID); err == nil { usage.Credits = amount.String() + } else if ctx.Err() == nil { + slog.Warn("读取生图积分用量失败", "job", job.ID, "operation", "image_charge", "error", err) } } return imageproxy.Task{ID: job.ID, UserID: job.UserID, Operation: job.Operation, Status: status, @@ -464,6 +492,9 @@ func (s *Service) recover(ctx context.Context) { } item, err := s.models.Get(ctx, job.ModelID) if err != nil { + if ctx.Err() == nil { + slog.Warn("加载待恢复生图模型失败", "job", job.ID, "operation", "recover_model", "error", err) + } continue } adapter := s.providers[item.Provider] @@ -474,7 +505,11 @@ func (s *Service) recover(ctx context.Context) { BaseURL: item.BaseURL, APIKey: item.APIKey, Protocol: string(item.Protocol)} callCtx, cancel := context.WithTimeout(ctx, time.Minute) result, err := adapter.TaskQuerier.QueryTask(callCtx, target, *job.ProviderJobID) - if err == nil && result.Status != "running" && result.Status != "pending" { + if err != nil { + if ctx.Err() == nil { + slog.Warn("查询待恢复生图任务失败", "job", job.ID, "operation", "query_task", "error_type", providerErrorType(err)) + } + } else if result.Status != "running" && result.Status != "pending" { s.handleResult(callCtx, job, result, nil) } cancel() diff --git a/monkeyai/backend/internal/imagegen/service_test.go b/monkeyai/backend/internal/imagegen/service_test.go index 75d6842ff..d66116083 100644 --- a/monkeyai/backend/internal/imagegen/service_test.go +++ b/monkeyai/backend/internal/imagegen/service_test.go @@ -4,8 +4,14 @@ import ( "bytes" "context" "errors" + "fmt" "image" "image/png" + "io" + "log/slog" + "net/http" + "net/url" + "strings" "sync" "testing" "time" @@ -298,3 +304,52 @@ func TestRejectedGenerationReleasesCredits(t *testing.T) { t.Fatalf("审核拒绝未退款: %s %+v", status, bills.finished) } } + +type testTransport func(*http.Request) (*http.Response, error) + +func (f testTransport) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +type testCloseBody struct { + io.Reader + err error +} + +func (b testCloseBody) Close() error { return b.err } + +func TestProviderErrorLogsDoNotExposeWrappedURL(t *testing.T) { + const secret = "private-token-123" + var logs bytes.Buffer + original := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(original) }) + + wrapped := fmt.Errorf("outer token=%s: %w", secret, &url.Error{ + Op: "Post", URL: "https://api.example/v1?token=" + secret, + Err: fmt.Errorf("inner token=%s: %w", secret, &url.Error{ + Op: "Get", URL: "https://api.example/v1?api_key=" + secret, + Err: errors.New("failure token=" + secret), + }), + }) + jobs := &testJobs{} + svc := NewService(nil, jobs, nil, nil, nil) + svc.handleResult(context.Background(), Job{ID: "job-1", Operation: "generate"}, ProviderResult{}, wrapped) + if jobs.job.Status != "unknown" || !strings.Contains(logs.String(), "error_type=url_error") || + strings.Contains(logs.String(), secret) || strings.Contains(logs.String(), "api_key=") { + t.Fatalf("上游错误日志包含敏感详情或状态错误: %s", logs.String()) + } + + logs.Reset() + client := &http.Client{Transport: testTransport(func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), + Body: testCloseBody{Reader: strings.NewReader(`{}`), err: wrapped}}, nil + })} + target := proxy.Target{BaseURL: "https://api.example/v1?token=" + secret, APIKey: secret, ModelID: "model-1"} + content, status, err := CallBody(context.Background(), client, target, "/images", []byte(`{}`), "application/json") + if err != nil || status != http.StatusOK || string(content) != "{}" { + t.Fatalf("上游响应异常: status=%d error=%v", status, err) + } + if !strings.Contains(logs.String(), "error_type=url_error") || !strings.Contains(logs.String(), "status=200") || + strings.Contains(logs.String(), secret) || strings.Contains(logs.String(), "token=") { + t.Fatalf("上游关闭日志包含敏感详情: %s", logs.String()) + } +} diff --git a/monkeyai/backend/internal/imagegen/upstream.go b/monkeyai/backend/internal/imagegen/upstream.go index 204f8c5f7..b3307c42b 100644 --- a/monkeyai/backend/internal/imagegen/upstream.go +++ b/monkeyai/backend/internal/imagegen/upstream.go @@ -7,6 +7,7 @@ import ( "encoding/json" "errors" "io" + "log/slog" "net/http" "net/url" "strings" @@ -48,7 +49,11 @@ func CallBody(ctx context.Context, client *http.Client, target proxy.Target, end if err != nil { return nil, 0, err } - defer response.Body.Close() + defer func() { + if closeErr := response.Body.Close(); closeErr != nil && ctx.Err() == nil { + slog.Warn("关闭生图上游响应失败", "operation", "close_upstream_response", "model_id", target.ModelID, "status", response.StatusCode, "error_type", providerErrorType(closeErr)) + } + }() if response.StatusCode < 200 || response.StatusCode >= 300 { return nil, response.StatusCode, nil } diff --git a/monkeyai/backend/internal/imageproxy/proxy.go b/monkeyai/backend/internal/imageproxy/proxy.go index c86b2da54..12fefe9b5 100644 --- a/monkeyai/backend/internal/imageproxy/proxy.go +++ b/monkeyai/backend/internal/imageproxy/proxy.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "io" + "log/slog" "net/http" "strings" "time" @@ -139,6 +140,9 @@ func (p *Proxy) upload(w http.ResponseWriter, r *http.Request) { } userID, err := p.keys.Authenticate(r.Context(), credential, "model:invoke") if err != nil || userID == "" { + if err != nil && r.Context().Err() == nil { + slog.Warn("生图鉴权失败", "operation", "authenticate_upload", "error", credentialError(err, credential)) + } unauthorized(w) return } @@ -182,6 +186,9 @@ func (p *Proxy) output(w http.ResponseWriter, r *http.Request) { } userID, err := p.keys.Authenticate(r.Context(), credential, "model:invoke") if err != nil || userID == "" { + if err != nil && r.Context().Err() == nil { + slog.Warn("生图鉴权失败", "operation", "authenticate_output", "error", credentialError(err, credential)) + } unauthorized(w) return } @@ -197,7 +204,9 @@ func (p *Proxy) output(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", mime) w.Header().Set("Cache-Control", "private, no-store") w.Header().Set("X-Content-Type-Options", "nosniff") - _, _ = w.Write(data) + if _, err := w.Write(data); err != nil && r.Context().Err() == nil { + slog.Warn("生图结果响应写入失败", "operation", "write_output", "user_id", userID, "error", err) + } } func (p *Proxy) generate(w http.ResponseWriter, r *http.Request) { @@ -288,6 +297,9 @@ func (p *Proxy) task(w http.ResponseWriter, r *http.Request) { } userID, err := p.keys.Authenticate(r.Context(), credential, "model:invoke") if err != nil || userID == "" { + if err != nil && r.Context().Err() == nil { + slog.Warn("生图鉴权失败", "operation", "authenticate_task", "error", credentialError(err, credential)) + } unauthorized(w) return } @@ -311,6 +323,9 @@ func (p *Proxy) resolve(w http.ResponseWriter, r *http.Request, credential, mode } target, err := p.resolver.Resolve(r.Context(), credential, model) if err != nil || target.UserID == "" { + if err != nil && r.Context().Err() == nil { + slog.Warn("生图模型解析失败", "operation", "resolve", "error", credentialError(err, credential)) + } unauthorized(w) return proxy.Target{}, false } @@ -356,6 +371,13 @@ func respond(w http.ResponseWriter, task Task, err error) { resource.JSON(w, http.StatusAccepted, task) } +func credentialError(err error, credential string) string { + if credential == "" { + return err.Error() + } + return strings.ReplaceAll(err.Error(), credential, "[redacted]") +} + func unauthorized(w http.ResponseWriter) { http.Error(w, http.StatusText(http.StatusUnauthorized), http.StatusUnauthorized) } diff --git a/monkeyai/backend/internal/imageproxy/proxy_test.go b/monkeyai/backend/internal/imageproxy/proxy_test.go index 507f555f0..0b3d9f6aa 100644 --- a/monkeyai/backend/internal/imageproxy/proxy_test.go +++ b/monkeyai/backend/internal/imageproxy/proxy_test.go @@ -96,7 +96,9 @@ func TestInputUploadAuthenticatesAndLimitsParts(t *testing.T) { t.Fatal(err) } } - writer.Close() + if err := writer.Close(); err != nil { + t.Fatal(err) + } r := httptest.NewRequest(http.MethodPost, "/v1/images/inputs", &body) r.Header.Set("Content-Type", writer.FormDataContentType()) if credential != "" { @@ -126,6 +128,16 @@ func (f outputFunc) Open(ctx context.Context, userID, id string) ([]byte, string return f(ctx, userID, id) } +type failedOutputWriter struct { + *httptest.ResponseRecorder + writes int +} + +func (w *failedOutputWriter) Write([]byte) (int, error) { + w.writes++ + return 0, errors.New("write failed") +} + func TestImageOutputRequiresSameKeyAndOwner(t *testing.T) { var data bytes.Buffer if err := png.Encode(&data, image.NewRGBA(image.Rect(0, 0, 2, 2))); err != nil { @@ -147,6 +159,13 @@ func TestImageOutputRequiresSameKeyAndOwner(t *testing.T) { if w.Code != http.StatusOK || w.Header().Get("Content-Type") != "image/png" || w.Header().Get("X-Content-Type-Options") != "nosniff" || !bytes.Equal(w.Body.Bytes(), data.Bytes()) { t.Fatalf("图片输出失败: %d, %s", w.Code, w.Body.String()) } + r := httptest.NewRequest(http.MethodGet, "/v1/images/outputs/output-1", nil) + r.Header.Set("Authorization", "Bearer invoke-key") + failed := &failedOutputWriter{ResponseRecorder: httptest.NewRecorder()} + h.ServeHTTP(failed, r) + if failed.Code != http.StatusOK || failed.writes != 1 { + t.Fatalf("写失败后发生了二次响应: status=%d writes=%d", failed.Code, failed.writes) + } for _, tc := range []struct { path, key string status int diff --git a/monkeyai/backend/internal/mcp/credential.go b/monkeyai/backend/internal/mcp/credential.go index 190403491..3cd587c5d 100644 --- a/monkeyai/backend/internal/mcp/credential.go +++ b/monkeyai/backend/internal/mcp/credential.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "io" + "log/slog" "net/http" "slices" "strings" @@ -34,7 +35,11 @@ func credentialStatus(c, cred resource.Object) string { return "authorization_required" } if c.String("authorization_method") == "http_header" { - b, _ := json.Marshal(cred["http_headers"]) + b, err := json.Marshal(cred["http_headers"]) + if err != nil { + slog.Error("编码认证 Header 失败", "connector_id", c.String("id"), "credential_id", cred.String("id"), "error", err) + return "authorization_required" + } if _, err := decodeHeaders(b); err == nil { return "authorized" } @@ -257,7 +262,7 @@ func (s *Service) saveCredential(w http.ResponseWriter, r *http.Request, admin b resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, chi.URLParam(r, "id"), chi.URLParam(r, "credentialID"), "save_credential") c, err := s.lockConnector(ctx, tx, chi.URLParam(r, "id"), u.ID, admin) if err == nil { err = manageCredential(c, admin) @@ -342,7 +347,11 @@ func (s *Service) saveCredential(w http.ResponseWriter, r *http.Request, admin b resource.Fail(w, resource.Invalid("请提供凭证名称或认证 Header")) return } - b, _ := json.Marshal(data) + b, err := json.Marshal(data) + if err != nil { + resource.Fail(w, fmt.Errorf("编码凭证更新: %w", err)) + return + } var out resource.Object if r.Method == http.MethodPost { out, err = resource.DecodeObject(queries.CreateCredential(ctx, b)) @@ -365,7 +374,9 @@ func (s *Service) saveCredential(w http.ResponseWriter, r *http.Request, admin b if in.Headers != nil { check, cancel := context.WithTimeout(context.WithoutCancel(ctx), time.Minute) defer cancel() - _, _ = s.testConnection(check, c, out, u.ID, admin) + if _, testErr := s.testConnection(check, c, out, u.ID, admin); testErr != nil { + slog.WarnContext(check, "保存凭证后连接测试失败", "connector_id", c.String("id"), "credential_id", id, "failure", safeMCPFailure(testErr)) + } out, err = s.Credential(check, s.Store.Pool, c, u.ID, id) } var views []resource.Object diff --git a/monkeyai/backend/internal/mcp/discovery.go b/monkeyai/backend/internal/mcp/discovery.go index 1588b0035..e68bf7472 100644 --- a/monkeyai/backend/internal/mcp/discovery.go +++ b/monkeyai/backend/internal/mcp/discovery.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "io" + "log/slog" "net/http" "net/url" "regexp" @@ -36,9 +37,12 @@ func discoveryURL(value, source string) bool { if !validURL(value) { return false } - u, _ := url.Parse(value) - origin, _ := url.Parse(source) - return origin != nil && (origin.Scheme != "https" || u.Scheme == "https") + u, err := url.Parse(value) + if err != nil { + return false + } + origin, err := url.Parse(source) + return err == nil && origin != nil && (origin.Scheme != "https" || u.Scheme == "https") } func readOAuthMetadata(ctx context.Context, h *http.Client, target string, out any) error { @@ -51,7 +55,11 @@ func readOAuthMetadata(ctx context.Context, h *http.Client, target string, out a if err != nil { return err } - defer resp.Body.Close() + defer func() { + if err := resp.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 OAuth 元数据响应失败", "operation", "read_metadata", "failure", safeMCPFailure(err)) + } + }() if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed { return metadataUnavailable } @@ -66,7 +74,10 @@ func readOAuthMetadata(ctx context.Context, h *http.Client, target string, out a } func metadataURLs(issuer, kind string) []string { - u, _ := url.Parse(issuer) + u, err := url.Parse(issuer) + if err != nil { + return nil + } origin := u.Scheme + "://" + u.Host path := u.EscapedPath() out := []string{origin + "/.well-known/" + kind + path} @@ -97,13 +108,18 @@ func discoverOAuth(ctx context.Context, target string) (oauthConfig, error) { defer h.CloseIdleConnections() // 不附带现有凭证;所有发现请求复用 MCP 的地址限制及禁止重定向策略。 - req, _ := http.NewRequestWithContext(ctx, http.MethodGet, target, nil) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil) + if err != nil { + return fail() + } req.Header.Set("Accept", "application/json, text/event-stream") resp, err := h.Do(req) if err != nil { return fail() } - resp.Body.Close() + if err := resp.Body.Close(); err != nil { + return fail() + } var metadataURL string if resp.StatusCode == http.StatusUnauthorized { for _, challenge := range resp.Header.Values("WWW-Authenticate") { @@ -136,7 +152,10 @@ func discoverOAuth(ctx context.Context, target string) (oauthConfig, error) { } } - u, _ := url.Parse(target) + u, err := url.Parse(target) + if err != nil { + return fail() + } issuer := u.Scheme + "://" + u.Host resourceURL := *u resourceURL.RawQuery, resourceURL.ForceQuery = "", false @@ -154,7 +173,10 @@ func discoverOAuth(ctx context.Context, target string) (oauthConfig, error) { if !discoveryURL(issuer, target) { return fail() } - issuerURL, _ := url.Parse(issuer) + issuerURL, err := url.Parse(issuer) + if err != nil { + return fail() + } if issuerURL.RawQuery != "" { return fail() } diff --git a/monkeyai/backend/internal/mcp/gateway.go b/monkeyai/backend/internal/mcp/gateway.go index 961f4b449..a46852fe3 100644 --- a/monkeyai/backend/internal/mcp/gateway.go +++ b/monkeyai/backend/internal/mcp/gateway.go @@ -5,6 +5,7 @@ import ( "context" "encoding/json" "errors" + "log/slog" "mime" "net/http" "net/url" @@ -34,22 +35,25 @@ func (s *Service) RegisterGateway(router chi.Router, keys KeyAuthenticator, bill w.Header().Set("Cache-Control", "no-store") if origin := r.Header.Get("Origin"); origin != "" { provided, err := url.Parse(origin) - public, _ := url.Parse(s.PublicURL) - if err != nil || public == nil || provided.User != nil || provided.Path != "" || provided.RawQuery != "" || provided.Fragment != "" || !strings.EqualFold(provided.Scheme, public.Scheme) || !strings.EqualFold(provided.Host, public.Host) { - rpcFail(w, nil, &resource.Error{Status: 403, Code: "invalid_origin", Message: "请求来源不被允许"}) + public, publicErr := url.Parse(s.PublicURL) + if publicErr != nil { + slog.ErrorContext(r.Context(), "MCP 公共地址配置无效", "operation", "origin_check", "failure_reason", "invalid_public_url") + } + if err != nil || publicErr != nil || public == nil || provided.User != nil || provided.Path != "" || provided.RawQuery != "" || provided.Fragment != "" || !strings.EqualFold(provided.Scheme, public.Scheme) || !strings.EqualFold(provided.Host, public.Host) { + rpcFail(r.Context(), w, nil, &resource.Error{Status: 403, Code: "invalid_origin", Message: "请求来源不被允许"}) return } } scheme, token, ok := strings.Cut(r.Header.Get("Authorization"), " ") if !ok || !strings.EqualFold(scheme, "Bearer") || strings.TrimSpace(token) == "" { w.Header().Set("WWW-Authenticate", `Bearer realm="mcp"`) - rpcFail(w, nil, &resource.Error{Status: 401, Code: "invalid_key", Message: "缺少 MCP 调用密钥"}) + rpcFail(r.Context(), w, nil, &resource.Error{Status: 401, Code: "invalid_key", Message: "缺少 MCP 调用密钥"}) return } user, err := keys.Authenticate(r.Context(), strings.TrimSpace(token), "mcp:invoke") if err != nil { w.Header().Set("WWW-Authenticate", `Bearer realm="mcp", error="invalid_token"`) - rpcFail(w, nil, &resource.Error{Status: 401, Code: "invalid_key", Message: "MCP 调用密钥无效或权限不足"}) + rpcFail(r.Context(), w, nil, &resource.Error{Status: 401, Code: "invalid_key", Message: "MCP 调用密钥无效或权限不足"}) return } if r.Method != http.MethodPost { @@ -59,11 +63,11 @@ func (s *Service) RegisterGateway(router chi.Router, keys KeyAuthenticator, bill } media, _, err := mime.ParseMediaType(r.Header.Get("Content-Type")) if err != nil || media != "application/json" { - rpcFail(w, nil, &resource.Error{Status: 415, Code: "invalid_content_type", Message: "请求须使用 application/json"}) + rpcFail(r.Context(), w, nil, &resource.Error{Status: 415, Code: "invalid_content_type", Message: "请求须使用 application/json"}) return } if version := r.Header.Get("MCP-Protocol-Version"); version != "" && !supportedVersion(version) { - rpcFail(w, nil, &resource.Error{Status: 400, Code: "unsupported_protocol", Message: "MCP 协议版本不支持"}) + rpcFail(r.Context(), w, nil, &resource.Error{Status: 400, Code: "unsupported_protocol", Message: "MCP 协议版本不支持"}) return } in, ok := readRequest(w, r) @@ -72,21 +76,21 @@ func (s *Service) RegisterGateway(router chi.Router, keys KeyAuthenticator, bill } connector, err := s.Connector(r.Context(), s.Store.Pool, chi.URLParam(r, "id"), user, false) if err != nil { - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } credentialID := chi.URLParam(r, "credentialID") if connector.String("authorization_mode") != "none" && credentialID == "" { - rpcFail(w, in.ID, selectionRequired) + rpcFail(r.Context(), w, in.ID, selectionRequired) return } credential, err := s.Credential(r.Context(), s.Store.Pool, connector, user, credentialID) if err != nil { - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } if connector.String("authorization_mode") != "none" && credentialStatus(connector, credential) != "authorized" { - rpcFail(w, in.ID, authorizationRequired) + rpcFail(r.Context(), w, in.ID, authorizationRequired) return } if strings.HasPrefix(in.Method, "notifications/") { @@ -130,12 +134,12 @@ func (s *Service) RegisterGateway(router chi.Router, keys KeyAuthenticator, bill func (s *Service) invoke(w http.ResponseWriter, r *http.Request, in request, connector, credential resource.Object, user string, billing InvocationBilling) { headers, credential, err := s.headers(r.Context(), connector, credential) if err != nil { - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } tools, err := credentialTools(r.Context(), s.Store.Pool, connector, credential, false) if err != nil { - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } if in.Method == "tools/list" { @@ -177,23 +181,25 @@ func (s *Service) invoke(w http.ResponseWriter, r *http.Request, in request, con } remote, err := openRemote(r.Context(), connector.String("url"), headers) if err != nil { - rpcFail(w, in.ID, &resource.Error{Status: 502, Code: "mcp_connect_failed", Message: "工具上游连接失败"}) + slog.WarnContext(r.Context(), "连接 MCP 上游失败", "connector_id", connector.String("id"), "credential_id", credential.String("id"), "operation", "connect", "failure", safeMCPFailure(err)) + rpcFail(r.Context(), w, in.ID, &resource.Error{Status: 502, Code: "mcp_connect_failed", Message: "工具上游连接失败"}) return } defer remote.close() id, err := billing.Begin(r.Context(), Invocation{UserID: user, ConnectorID: connector.String("id"), CredentialID: credential.String("id"), ToolID: tool.String("id"), SessionID: r.Header.Get("X-Session-ID"), IdempotencyKey: r.Header.Get("Idempotency-Key"), RequestHash: resource.Hash(resource.Object{"connector": connector.String("id"), "credential": credential.String("id"), "params": params, "session_id": r.Header.Get("X-Session-ID")})}) if err != nil { - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } w.Header().Set("X-Billing-Transaction-ID", id) if err = billing.Start(r.Context(), id); err != nil { ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 30*time.Second) defer cancel() - if billing.Finish(ctx, id, InvocationResult{Known: true, Result: "failed", ErrorCode: "mcp_not_started"}) != nil { + if finishErr := billing.Finish(ctx, id, InvocationResult{Known: true, Result: "failed", ErrorCode: "mcp_not_started"}); finishErr != nil { + slog.ErrorContext(ctx, "MCP 调用开始失败后结算失败", "connector_id", connector.String("id"), "operation", "finish", "failure", safeMCPFailure(finishErr)) w.Header().Set("X-Billing-Status", "pending") } - rpcFail(w, in.ID, err) + rpcFail(r.Context(), w, in.ID, err) return } result, callErr := remote.call(r.Context(), 2, "tools/call", in.Params) @@ -221,9 +227,13 @@ func (s *Service) invoke(w http.ResponseWriter, r *http.Request, in request, con ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 30*time.Second) defer cancel() if err = billing.Finish(ctx, id, outcome); err != nil { + slog.ErrorContext(ctx, "MCP 调用结算失败", "connector_id", connector.String("id"), "operation", "finish", "failure", safeMCPFailure(err)) w.Header().Set("X-Billing-Status", "pending") } if callErr != nil { + if rpc == nil { + slog.WarnContext(r.Context(), "MCP 工具调用失败", "connector_id", connector.String("id"), "credential_id", credential.String("id"), "operation", "tools/call", "failure", safeMCPFailure(callErr)) + } code := -32603 if rpc != nil { code = rpc.Code diff --git a/monkeyai/backend/internal/mcp/icon.go b/monkeyai/backend/internal/mcp/icon.go index 6d2931420..9d95505a8 100644 --- a/monkeyai/backend/internal/mcp/icon.go +++ b/monkeyai/backend/internal/mcp/icon.go @@ -2,19 +2,21 @@ package mcp import ( "bytes" + "errors" "fmt" "image" _ "image/jpeg" _ "image/png" "io" + "log/slog" "net/http" "strings" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/mcp/sqlc" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/resource" - "github.com/go-chi/chi/v5" + "github.com/jackc/pgx/v5" ) func (s *Service) WithStorage(storage resource.Storage) *Service { s.storage = storage; return s } @@ -41,13 +43,21 @@ func (s *Service) uploadIcon(w http.ResponseWriter, r *http.Request, admin bool) resource.Fail(w, resource.Invalid("图标必须小于 1 MiB")) return } - defer r.MultipartForm.RemoveAll() + defer func() { + if err := r.MultipartForm.RemoveAll(); err != nil { + slog.WarnContext(r.Context(), "清理图标上传临时文件失败", "connector_id", chi.URLParam(r, "id"), "error", err) + } + }() f, _, err := r.FormFile("icon") if err != nil { resource.Fail(w, resource.Invalid("缺少 icon 文件")) return } - defer f.Close() + defer func() { + if err := f.Close(); err != nil { + slog.WarnContext(r.Context(), "关闭 Connector 图标上传文件失败", "connector_id", chi.URLParam(r, "id"), "operation", "close_upload", "error", err) + } + }() data, err := io.ReadAll(io.LimitReader(f, (1<<20)+1)) if err != nil || len(data) > 1<<20 { resource.Fail(w, resource.Invalid("图标超限")) @@ -64,7 +74,7 @@ func (s *Service) uploadIcon(w http.ResponseWriter, r *http.Request, admin bool) resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, chi.URLParam(r, "id"), "", "upload_icon") id := chi.URLParam(r, "id") u, _ := identity.UserFromContext(ctx) user := u.ID @@ -115,6 +125,9 @@ func (s *Service) uploadIcon(w http.ResponseWriter, r *http.Request, admin bool) func (s *Service) icon(w http.ResponseWriter, r *http.Request, connector string) { key, err := sqlc.New(s.Store.Pool).GetConnectorIcon(r.Context(), connector) if err != nil || key == "" { + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(r.Context(), "读取 Connector 图标失败", "connector_id", connector, "error", err) + } resource.Fail(w, resource.NotFound) return } @@ -124,7 +137,11 @@ func (s *Service) icon(w http.ResponseWriter, r *http.Request, connector string) resource.Fail(w, err) return } - defer body.Close() + defer func() { + if err := body.Close(); err != nil { + slog.WarnContext(r.Context(), "关闭 Connector 图标读取流失败", "connector_id", connector, "operation", "close_icon", "failure", safeMCPFailure(err)) + } + }() mime := "image/png" if strings.HasSuffix(key, ".jpeg") { mime = "image/jpeg" @@ -133,5 +150,7 @@ func (s *Service) icon(w http.ResponseWriter, r *http.Request, connector string) w.Header().Set("X-Content-Type-Options", "nosniff") w.Header().Set("Cache-Control", "private, no-cache") w.Header().Set("ETag", `"`+resource.Hash(key)+`"`) - _, _ = io.Copy(w, body) + if _, err := io.Copy(w, body); err != nil { + slog.WarnContext(r.Context(), "传输 Connector 图标失败", "connector_id", connector, "error", err) + } } diff --git a/monkeyai/backend/internal/mcp/oauth.go b/monkeyai/backend/internal/mcp/oauth.go index 3bb7df98c..329e42b49 100644 --- a/monkeyai/backend/internal/mcp/oauth.go +++ b/monkeyai/backend/internal/mcp/oauth.go @@ -11,6 +11,7 @@ import ( "fmt" "io" "log/slog" + "net" "net/http" "net/url" "strconv" @@ -23,6 +24,7 @@ import ( "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/resource" "github.com/go-chi/chi/v5" "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" ) type oauthConfig struct { @@ -38,15 +40,30 @@ type oauthConfig struct { } func oauthSettings(c resource.Object) oauthConfig { - b, _ := json.Marshal(c["oauth_config"]) + b, err := json.Marshal(c["oauth_config"]) + if err != nil { + var unsupported *json.UnsupportedTypeError + if errors.As(err, &unsupported) { + slog.Error("编码 OAuth 配置失败", "connector_id", c.String("id"), "operation", "encode_config", "error", unsupported) + } else { + // 自定义 JSON 错误可能回显配置中的客户端密钥。 + slog.Error("编码 OAuth 配置失败", "connector_id", c.String("id"), "operation", "encode_config", "failure_reason", "invalid_config_encoding") + } + return oauthConfig{} + } var o oauthConfig - _ = json.Unmarshal(b, &o) + if err := json.Unmarshal(b, &o); err != nil { + slog.Error("解析 OAuth 配置失败", "connector_id", c.String("id"), "operation", "decode_config", "failure_reason", "invalid_config_format") + return oauthConfig{} + } return o } -func token() string { +func token() (string, error) { b := make([]byte, 32) - _, _ = rand.Read(b) - return base64.RawURLEncoding.EncodeToString(b) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("生成 OAuth 随机数: %w", err) + } + return base64.RawURLEncoding.EncodeToString(b), nil } func hash(value string) string { v := sha256.Sum256([]byte(value)); return hex.EncodeToString(v[:]) } func (s *Service) callbackURL(id string) string { @@ -60,7 +77,7 @@ func (s *Service) authorize(w http.ResponseWriter, r *http.Request, admin bool) resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, chi.URLParam(r, "id"), chi.URLParam(r, "credentialID"), "authorize") c, err := s.lockConnector(ctx, tx, chi.URLParam(r, "id"), u.ID, admin) if err == nil { err = manageCredential(c, admin) @@ -112,7 +129,16 @@ func (s *Service) authorize(w http.ResponseWriter, r *http.Request, admin bool) } } } - state, verifier := token(), token() + state, err := token() + if err != nil { + resource.Fail(w, err) + return + } + verifier, err := token() + if err != nil { + resource.Fail(w, err) + return + } redirect := s.callbackURL(c.String("id")) if err = s.ensureOAuthClient(ctx, tx, c, redirect); err != nil { resource.Fail(w, err) @@ -120,7 +146,11 @@ func (s *Service) authorize(w http.ResponseWriter, r *http.Request, admin bool) } data["config_revision"] = c.Int("config_revision") data["state_hash"], data["verifier"], data["redirect_uri"] = hash(state), verifier, redirect - b, _ := json.Marshal(data) + b, err := json.Marshal(data) + if err != nil { + resource.Fail(w, fmt.Errorf("编码 OAuth 授权请求: %w", err)) + return + } request, err := resource.DecodeObject(sqlc.New(tx).CreateOAuthRequest(ctx, b)) if err == nil { err = tx.Commit(ctx) @@ -130,7 +160,11 @@ func (s *Service) authorize(w http.ResponseWriter, r *http.Request, admin bool) return } o := oauthSettings(c) - target, _ := url.Parse(o.AuthorizationURL) + target, err := url.Parse(o.AuthorizationURL) + if err != nil || target == nil || target.Scheme == "" || target.Host == "" { + resource.Fail(w, resource.Invalid("OAuth 授权地址无效")) + return + } q := target.Query() q.Set("response_type", "code") q.Set("client_id", o.ClientID) @@ -198,6 +232,9 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { } request, err := resource.DecodeObject(sqlc.New(s.Store.Pool).ConsumeOAuthRequest(ctx, sqlc.ConsumeOAuthRequestParams{StateHash: hash(state), ConnectorID: id})) if err != nil { + if !errors.Is(err, pgx.ErrNoRows) { + slog.ErrorContext(ctx, "消费 OAuth 授权事务失败", "connector_id", id, "operation", "consume", "failure", safeMCPFailure(err)) + } resource.Fail(w, resource.Invalid("授权事务无效或已使用")) return } @@ -208,7 +245,9 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { } cleanup, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) defer cancel() - _, _ = sqlc.New(s.Store.Pool).FinishOAuthRequest(cleanup, sqlc.FinishOAuthRequestParams{ID: request.String("id"), Status: "failed", CredentialID: ""}) + if _, err := sqlc.New(s.Store.Pool).FinishOAuthRequest(cleanup, sqlc.FinishOAuthRequestParams{ID: request.String("id"), Status: "failed", CredentialID: ""}); err != nil { + slog.ErrorContext(cleanup, "标记 OAuth 授权失败事务失败", "connector_id", id, "operation", "finish_failed", "error", err) + } }() if r.URL.Query().Get("error") != "" || r.URL.Query().Get("code") == "" { resource.Fail(w, resource.Invalid("授权已取消或缺少授权码")) @@ -220,13 +259,20 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { return } c, err := s.oauthContext(ctx, tx, request) - _ = tx.Rollback(ctx) + if rollbackErr := tx.Rollback(ctx); rollbackErr != nil { + if err == nil { + err = rollbackErr + } else { + slog.WarnContext(ctx, "释放 OAuth 授权事务失败", "connector_id", id, "operation", "rollback", "error", rollbackErr) + } + } if err != nil { resource.Fail(w, err) return } result, err := exchange(ctx, c, url.Values{"grant_type": {"authorization_code"}, "code": {r.URL.Query().Get("code")}, "redirect_uri": {request.String("redirect_uri")}, "code_verifier": {request.String("verifier")}}) if err != nil { + slog.WarnContext(ctx, "OAuth Token 交换失败", "connector_id", id, "operation", "exchange", "failure", safeMCPFailure(err)) resource.Fail(w, &resource.Error{Status: 502, Code: "oauth_exchange_failed", Message: "OAuth Token 交换失败,请重新发起授权"}) return } @@ -235,7 +281,7 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, id, request.String("credential_id"), "callback") c, err = s.oauthContext(ctx, tx, request) if err != nil { resource.Fail(w, err) @@ -253,7 +299,11 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { data := resource.Object{"id": credential, "connector_id": id, "user_id": user, "name": request.String("name"), "http_headers": resource.Object{}, "oauth_access_token": result.Access, "oauth_refresh_token": result.Refresh, "oauth_expires_at": result.Expires, "config_revision": c.Int("config_revision"), "auth_change": true} - b, _ := json.Marshal(data) + b, err := json.Marshal(data) + if err != nil { + resource.Fail(w, fmt.Errorf("编码 OAuth 凭证: %w", err)) + return + } queries := sqlc.New(tx) var cred resource.Object if create { @@ -284,6 +334,9 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { check, cancel := context.WithTimeout(context.WithoutCancel(ctx), time.Minute) defer cancel() _, testErr := s.testConnection(check, c, cred, request.String("user_id"), c.String("authorization_mode") == "centralized") + if testErr != nil { + slog.WarnContext(check, "OAuth 授权后连接测试失败", "connector_id", id, "credential_id", credential, "failure", safeMCPFailure(testErr)) + } finish, stop := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) defer stop() rows, err := sqlc.New(s.Store.Pool).FinishOAuthRequest(finish, sqlc.FinishOAuthRequestParams{ID: request.String("id"), Status: "succeeded", CredentialID: credential}) @@ -300,7 +353,9 @@ func (s *Service) Callback(w http.ResponseWriter, r *http.Request) { message = "授权成功,但自动连接测试失败,请返回 MonkeyAI 查看凭证状态并重试连接测试。" } w.Header().Set("Content-Type", "text/html; charset=utf-8") - _, _ = fmt.Fprintf(w, "
%s
", message) + if _, err := fmt.Fprintf(w, "%s
", message); err != nil { + slog.WarnContext(ctx, "写入 OAuth 授权结果失败", "connector_id", id, "credential_id", credential, "failure_reason", "response_write_failed", "error_type", fmt.Sprintf("%T", err)) + } } type tokens struct { @@ -310,6 +365,54 @@ type tokens struct { var invalidGrant = errors.New("OAuth 凭证已失效") +type tokenExchangeError struct { + reason string + status int +} + +func (e tokenExchangeError) Error() string { return e.reason } + +// 上游及回调错误可能携带 URL 查询串、授权码或 Token,仅输出受控分类与状态。 +func safeMCPFailure(err error) []any { + var exchangeErr tokenExchangeError + if errors.As(err, &exchangeErr) { + if exchangeErr.status != 0 { + return []any{"reason", exchangeErr.reason, "upstream_status", exchangeErr.status} + } + return []any{"reason", exchangeErr.reason} + } + var failure *resource.Error + if errors.As(err, &failure) { + return []any{"reason", failure.Code, "status", failure.Status} + } + var upstream remoteStatus + if errors.As(err, &upstream) { + return []any{"reason", "upstream_http_error", "upstream_status", int(upstream)} + } + if errors.Is(err, context.DeadlineExceeded) { + return []any{"reason", "timeout"} + } + if errors.Is(err, context.Canceled) { + return []any{"reason", "canceled"} + } + var network net.Error + if errors.As(err, &network) { + return []any{"reason", "network_error", "timeout", network.Timeout()} + } + var urlErr *url.Error + if errors.As(err, &urlErr) { + return []any{"reason", "transport_error"} + } + var databaseErr *pgconn.PgError + if errors.As(err, &databaseErr) { + return []any{"reason", "database_error", "sqlstate", databaseErr.Code} + } + if errors.Is(err, pgx.ErrNoRows) { + return []any{"error", pgx.ErrNoRows} + } + return []any{"reason", "internal_error", "error_type", fmt.Sprintf("%T", err)} +} + func exchange(ctx context.Context, c resource.Object, v url.Values) (tokens, error) { o := oauthSettings(c) if o.clientSecretExpired() { @@ -325,7 +428,7 @@ func exchange(ctx context.Context, c resource.Object, v url.Values) (tokens, err } req, err := http.NewRequestWithContext(ctx, "POST", o.TokenURL, strings.NewReader(v.Encode())) if err != nil { - return tokens{}, err + return tokens{}, tokenExchangeError{reason: "invalid_token_url"} } if o.TokenAuthMethod == "client_secret_basic" { req.SetBasicAuth(url.QueryEscape(o.ClientID), url.QueryEscape(secret)) @@ -338,10 +441,17 @@ func exchange(ctx context.Context, c resource.Object, v url.Values) (tokens, err if err != nil { return tokens{}, err } - defer resp.Body.Close() + defer func() { + if err := resp.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 OAuth Token 响应失败", "operation", "exchange", "failure", safeMCPFailure(err)) + } + }() data, err := io.ReadAll(io.LimitReader(resp.Body, (1<<20)+1)) - if err != nil || len(data) > 1<<20 { - return tokens{}, fmt.Errorf("Token 响应无效") + if err != nil { + return tokens{}, tokenExchangeError{reason: "response_read_error", status: resp.StatusCode} + } + if len(data) > 1<<20 { + return tokens{}, tokenExchangeError{reason: "response_too_large", status: resp.StatusCode} } var payload struct { Access string `json:"access_token"` @@ -350,22 +460,34 @@ func exchange(ctx context.Context, c resource.Object, v url.Values) (tokens, err Type string `json:"token_type"` Error string `json:"error"` } + var form url.Values if json.Unmarshal(data, &payload) != nil { - form, err := url.ParseQuery(string(data)) + form, err = url.ParseQuery(string(data)) if err != nil { - return tokens{}, fmt.Errorf("Token 响应无效") + return tokens{}, tokenExchangeError{reason: "invalid_token_response", status: resp.StatusCode} } payload.Access = form.Get("access_token") payload.Refresh = form.Get("refresh_token") payload.Type = form.Get("token_type") payload.Error = form.Get("error") - payload.Expires, _ = strconv.ParseInt(form.Get("expires_in"), 10, 64) } if resp.StatusCode == 400 && payload.Error == "invalid_grant" { return tokens{}, invalidGrant } - if resp.StatusCode != 200 || payload.Error != "" || payload.Access == "" || strings.ContainsAny(payload.Access, "\r\n") || (!strings.EqualFold(payload.Type, "bearer") && payload.Type != "") { - return tokens{}, fmt.Errorf("Token 交换失败") + if resp.StatusCode != http.StatusOK { + return tokens{}, tokenExchangeError{reason: "upstream_http_error", status: resp.StatusCode} + } + if payload.Error != "" || payload.Access == "" || strings.ContainsAny(payload.Access, "\r\n") || (!strings.EqualFold(payload.Type, "bearer") && payload.Type != "") { + return tokens{}, tokenExchangeError{reason: "invalid_token_response", status: resp.StatusCode} + } + if form != nil { + if expires := form.Get("expires_in"); expires != "" { + seconds, err := strconv.ParseInt(expires, 10, 64) + if err != nil { + return tokens{}, tokenExchangeError{reason: "invalid_expires_in", status: resp.StatusCode} + } + payload.Expires = seconds + } } result := tokens{Access: payload.Access, Refresh: payload.Refresh} if payload.Expires > 0 { @@ -379,7 +501,7 @@ func (s *Service) refresh(ctx context.Context, c, cred resource.Object) (resourc if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, c.String("id"), cred.String("id"), "refresh") current, err := resource.DecodeObject(sqlc.New(tx).LockConnector(ctx, c.String("id"))) if err != nil { return nil, err @@ -417,6 +539,7 @@ func (s *Service) refresh(ctx context.Context, c, cred resource.Object) (resourc return nil, authorizationRequired } if err != nil { + slog.WarnContext(ctx, "刷新 OAuth 凭证失败", "connector_id", c.String("id"), "credential_id", cred.String("id"), "operation", "refresh", "failure", safeMCPFailure(err)) return nil, &resource.Error{Status: 502, Code: "oauth_refresh_failed", Message: "OAuth 刷新暂时失败,请稍后重试"} } if result.Refresh == "" { @@ -453,13 +576,15 @@ func (s *Service) refreshCredentials(ctx context.Context) { defer stop() c, err := resource.DecodeObject(item.Connector, nil) if err != nil { + slog.ErrorContext(refresh, "解析待刷新 Connector 数据失败", "operation", "decode_connector", "error", err) return } cred, err := resource.DecodeObject(item.Credential, nil) if err != nil { + slog.ErrorContext(refresh, "解析待刷新凭证数据失败", "connector_id", c.String("id"), "operation", "decode_credential", "error", err) return } - if _, err = s.refresh(refresh, c, cred); err != nil && ctx.Err() == nil { + if _, err = s.refresh(refresh, c, cred); err != nil && ctx.Err() == nil && !errors.Is(err, authorizationRequired) && !isOAuthRefreshFailure(err) { slog.WarnContext(ctx, "Connector OAuth 自动刷新失败", "connector_id", c.String("id"), "credential_id", cred.String("id"), "error", err) } }) @@ -471,7 +596,9 @@ func (s *Service) Run(ctx context.Context) { defer ticker.Stop() for { cleanup, cancel := context.WithTimeout(ctx, 10*time.Second) - _ = sqlc.New(s.Store.Pool).CleanupOAuthRequests(cleanup) + if err := sqlc.New(s.Store.Pool).CleanupOAuthRequests(cleanup); err != nil && ctx.Err() == nil { + slog.ErrorContext(cleanup, "清理过期 OAuth 授权请求失败", "operation", "cleanup", "error", err) + } cancel() s.refreshCredentials(ctx) select { @@ -481,3 +608,8 @@ func (s *Service) Run(ctx context.Context) { } } } + +func isOAuthRefreshFailure(err error) bool { + var failure *resource.Error + return errors.As(err, &failure) && failure.Code == "oauth_refresh_failed" +} diff --git a/monkeyai/backend/internal/mcp/oauth_test.go b/monkeyai/backend/internal/mcp/oauth_test.go index 83a8919f6..6234a28fb 100644 --- a/monkeyai/backend/internal/mcp/oauth_test.go +++ b/monkeyai/backend/internal/mcp/oauth_test.go @@ -3,9 +3,12 @@ package mcp import ( "context" "encoding/json" + "errors" + "fmt" "net/http" "net/http/httptest" "net/url" + "strings" "sync" "sync/atomic" "testing" @@ -14,6 +17,73 @@ import ( "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/resource" ) +func TestExchangeRejectsInvalidExpiresIn(t *testing.T) { + t.Setenv("MONKEYAI_MCP_ALLOWED_CIDRS", "127.0.0.0/8") + for _, expiry := range []string{"invalid", "999999999999999999999999"} { + t.Run(expiry, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/x-www-form-urlencoded") + fmt.Fprint(w, "access_token=private-token&token_type=bearer&expires_in="+expiry) + })) + defer server.Close() + _, err := exchange(t.Context(), resource.Object{"oauth_config": oauthConfig{ClientID: "client", TokenURL: server.URL}}, url.Values{"grant_type": {"authorization_code"}}) + var failure tokenExchangeError + if !errors.As(err, &failure) || failure.reason != "invalid_expires_in" || failure.status != http.StatusOK { + t.Fatalf("无效 expires_in 未被识别: %v", err) + } + if strings.Contains(fmt.Sprint(safeMCPFailure(err)), "private-token") { + t.Fatal("日志分类包含访问令牌") + } + }) + } +} + +func TestExchangeHTTPFailureExcludesResponse(t *testing.T) { + t.Setenv("MONKEYAI_MCP_ALLOWED_CIDRS", "127.0.0.0/8") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusBadGateway) + fmt.Fprint(w, "access_token=private-token&error=upstream-failed") + })) + defer server.Close() + _, err := exchange(t.Context(), resource.Object{"oauth_config": oauthConfig{ClientID: "client", TokenURL: server.URL}}, url.Values{"grant_type": {"authorization_code"}}) + logged := fmt.Sprint(safeMCPFailure(err)) + if strings.Contains(logged, "private-token") || !strings.Contains(logged, "502") || !strings.Contains(logged, "upstream_http_error") { + t.Fatalf("上游错误状态记录不安全: %s", logged) + } +} + +func TestOAuthFailureDoesNotExposeURL(t *testing.T) { + err := &url.Error{Op: "POST", URL: "https://oauth.example/token?code=private-code", Err: errors.New("private-token")} + logged := fmt.Sprint(safeMCPFailure(err)) + if strings.Contains(logged, "private-code") || strings.Contains(logged, "private-token") || !strings.Contains(logged, "network_error") { + t.Fatalf("上游错误分类不安全: %s", logged) + } +} + +func TestAutomaticRefreshContinuesAfterFailure(t *testing.T) { + f := setup(t) + remote := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Error(err) + } + if r.Form.Get("refresh_token") == "broken" { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + resource.JSON(w, http.StatusOK, resource.Object{"access_token": "renewed", "expires_in": 3600}) + })) + defer remote.Close() + c := f.call("POST", "/agent/connectors", resource.Object{"name": "刷新任务隔离", "url": remote.URL, "authorization_mode": "independent", "authorization_method": "oauth", "oauth_config": resource.Object{"client_id": "client", "token_url": remote.URL, "authorization_url": remote.URL}}, "owner", "", 201) + for _, refresh := range []string{"broken", "valid"} { + f.sql(`INSERT INTO connector_credentials(id,connector_id,user_id,name,oauth_access_token,oauth_refresh_token,oauth_expires_at,config_revision) VALUES($1,$2,$3,$4,'old',$5,now()-interval '1 minute',1)`, resource.ID(), c.String("id"), f.users["owner"], refresh, refresh) + } + f.service.refreshCredentials(t.Context()) + var access string + if err := f.pool.QueryRow(t.Context(), `SELECT oauth_access_token FROM connector_credentials WHERE connector_id=$1 AND name='valid'`, c.String("id")).Scan(&access); err != nil || access != "renewed" { + t.Fatalf("单条刷新失败阻断了后续凭证: %s, %v", access, err) + } +} + func TestAutomaticRefresh(t *testing.T) { f := setup(t) var calls atomic.Int32 diff --git a/monkeyai/backend/internal/mcp/protocol.go b/monkeyai/backend/internal/mcp/protocol.go index 6fbc1ab96..fe6ce2077 100644 --- a/monkeyai/backend/internal/mcp/protocol.go +++ b/monkeyai/backend/internal/mcp/protocol.go @@ -2,9 +2,11 @@ package mcp import ( "bytes" + "context" "encoding/json" "errors" "io" + "log/slog" "net/http" "slices" @@ -84,7 +86,7 @@ func rpcReply(w http.ResponseWriter, status int, id json.RawMessage, result any, resource.JSON(w, status, out) } -func rpcFail(w http.ResponseWriter, id json.RawMessage, err error) { +func rpcFail(ctx context.Context, w http.ResponseWriter, id json.RawMessage, err error) { var failure *resource.Error var postgres *pgconn.PgError if !errors.As(err, &failure) { @@ -94,12 +96,17 @@ func rpcFail(w http.ResponseWriter, id json.RawMessage, err error) { case errors.As(err, &postgres) && postgres.Code == "22P02": failure = &resource.Error{Status: 400, Code: "invalid_request", Message: "资源标识无效"} default: + slog.ErrorContext(ctx, "MCP 请求处理失败", "operation", "rpc", "failure", safeMCPFailure(err)) failure = &resource.Error{Status: 500, Code: "mcp_internal_error", Message: "MCP 请求处理失败"} } } if references, ok := failure.References.(map[string]string); ok && references["transaction_id"] != "" { w.Header().Set("X-Billing-Transaction-ID", references["transaction_id"]) } - data, _ := json.Marshal(resource.Object{"code": failure.Code, "references": failure.References}) + data, marshalErr := json.Marshal(resource.Object{"code": failure.Code, "references": failure.References}) + if marshalErr != nil { + slog.Error("编码 MCP 错误响应失败", "code", failure.Code, "error", marshalErr) + data = nil + } rpcReply(w, failure.Status, id, nil, &rpcError{Code: -32000, Message: failure.Message, Data: data}) } diff --git a/monkeyai/backend/internal/mcp/registration.go b/monkeyai/backend/internal/mcp/registration.go index 45bfc0d26..17f7ec83a 100644 --- a/monkeyai/backend/internal/mcp/registration.go +++ b/monkeyai/backend/internal/mcp/registration.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io" + "log/slog" "net/http" "strings" "time" @@ -26,11 +27,13 @@ func (s *Service) ensureOAuthClient(ctx context.Context, tx pgx.Tx, c resource.O var err error o, err = discoverOAuth(ctx, c.String("url")) if err != nil { + slog.WarnContext(ctx, "OAuth 自动发现失败", "connector_id", c.String("id"), "operation", "discovery", "error", err) return &resource.Error{Status: 502, Code: "oauth_discovery_failed", Message: "OAuth 自动发现失败,请确认 MCP 服务支持元数据发现和动态客户端注册,或切换为手动配置"} } } registered, err := registerOAuthClient(ctx, o, redirect) if err != nil { + slog.WarnContext(ctx, "OAuth 客户端注册失败", "connector_id", c.String("id"), "operation", "registration", "error", err) return &resource.Error{Status: 502, Code: "oauth_registration_failed", Message: "OAuth 动态客户端注册失败,请检查注册端点或手动配置 Client ID"} } revision := c.Int("config_revision") @@ -42,7 +45,10 @@ func (s *Service) ensureOAuthClient(ctx context.Context, tx pgx.Tx, c resource.O } } o.ClientID, o.TokenAuthMethod, o.ClientSecretExpiresAt = registered.ClientID, registered.Method, registered.SecretExpires - data, _ := json.Marshal(resource.Object{"id": c.String("id"), "oauth_config": o, "oauth_client_secret": registered.Secret, "config_revision": revision}) + data, err := json.Marshal(resource.Object{"id": c.String("id"), "oauth_config": o, "oauth_client_secret": registered.Secret, "config_revision": revision}) + if err != nil { + return fmt.Errorf("编码 OAuth 客户端配置: %w", err) + } if err = connector.New(tx).UpdateResource(ctx, data); err != nil { return err } @@ -82,7 +88,10 @@ func registerOAuthClient(ctx context.Context, o oauthConfig, redirect string) (o if o.Scopes != "" { body["scope"] = o.Scopes } - data, _ := json.Marshal(body) + data, err := json.Marshal(body) + if err != nil { + return oauthRegistration{}, fmt.Errorf("编码 OAuth 注册请求: %w", err) + } req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.RegistrationURL, bytes.NewReader(data)) if err != nil { return fail() @@ -95,7 +104,11 @@ func registerOAuthClient(ctx context.Context, o oauthConfig, redirect string) (o if err != nil { return fail() } - defer resp.Body.Close() + defer func() { + if err := resp.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 OAuth 注册响应失败", "operation", "register_client", "failure", safeMCPFailure(err)) + } + }() if resp.StatusCode != http.StatusCreated { return fail() } diff --git a/monkeyai/backend/internal/mcp/remote.go b/monkeyai/backend/internal/mcp/remote.go index 475c1a407..f701088e8 100644 --- a/monkeyai/backend/internal/mcp/remote.go +++ b/monkeyai/backend/internal/mcp/remote.go @@ -7,6 +7,7 @@ import ( "encoding/json" "fmt" "io" + "log/slog" "mime" "net/http" "strings" @@ -54,8 +55,12 @@ func (c *remoteClient) close() { ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() response, err := c.request(ctx, http.MethodDelete, nil) - if err == nil { - response.Body.Close() + if err != nil { + slog.WarnContext(ctx, "关闭 MCP 会话失败", "operation", "delete_session", "failure", safeMCPFailure(err)) + return + } + if err := response.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 MCP 会话响应失败", "operation", "delete_session", "failure", safeMCPFailure(err)) } } @@ -68,7 +73,11 @@ func (c *remoteClient) call(ctx context.Context, id int, method string, params a if err != nil { return nil, err } - defer response.Body.Close() + defer func() { + if err := response.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 MCP 调用响应失败", "operation", "call", "failure", safeMCPFailure(err)) + } + }() if response.StatusCode < 200 || response.StatusCode >= 300 { return nil, remoteStatus(response.StatusCode) } @@ -164,7 +173,9 @@ func openRemote(ctx context.Context, target string, headers map[string]string) ( if err != nil { return nil, err } - response.Body.Close() + if err := response.Body.Close(); err != nil { + slog.WarnContext(ctx, "关闭 MCP 握手响应失败", "operation", "initialize", "failure", safeMCPFailure(err)) + } if response.StatusCode < 200 || response.StatusCode >= 300 { return nil, remoteStatus(response.StatusCode) } diff --git a/monkeyai/backend/internal/mcp/service.go b/monkeyai/backend/internal/mcp/service.go index 1ff7e3b45..a37eaf87d 100644 --- a/monkeyai/backend/internal/mcp/service.go +++ b/monkeyai/backend/internal/mcp/service.go @@ -2,6 +2,8 @@ package mcp import ( "context" + "errors" + "log/slog" "net/http" "strings" @@ -257,7 +259,7 @@ func (s *Service) updateTool(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer rollbackMCP(r.Context(), tx, chi.URLParam(r, "id"), "", "update_tool") var mode string mode, err = sqlc.New(tx).GetAuthorizationMode(r.Context(), chi.URLParam(r, "id")) if err != nil { @@ -292,3 +294,9 @@ func (s *Service) updateTool(w http.ResponseWriter, r *http.Request) { resource.JSON(w, 200, o) } + +func rollbackMCP(ctx context.Context, tx pgx.Tx, connectorID, credentialID, operation string) { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) { + slog.WarnContext(ctx, "回滚 MCP 事务失败", "operation", operation, "connector_id", connectorID, "credential_id", credentialID, "failure", safeMCPFailure(err)) + } +} diff --git a/monkeyai/backend/internal/mcp/tools.go b/monkeyai/backend/internal/mcp/tools.go index 34d7e89cb..a3e6586f6 100644 --- a/monkeyai/backend/internal/mcp/tools.go +++ b/monkeyai/backend/internal/mcp/tools.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity" @@ -93,7 +94,10 @@ func (s *Service) headers(ctx context.Context, c, cred resource.Object) (map[str } return map[string]string{"Authorization": "Bearer " + fresh.String("oauth_access_token")}, fresh, nil } - b, _ := json.Marshal(cred["http_headers"]) + b, err := json.Marshal(cred["http_headers"]) + if err != nil { + return nil, nil, fmt.Errorf("编码认证 Header: %w", err) + } headers, err := decodeHeaders(b) return headers, cred, err } @@ -135,7 +139,7 @@ func (s *Service) testConnection(ctx context.Context, c, cred resource.Object, u if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer rollbackMCP(ctx, tx, c.String("id"), cred.String("id"), "test_connection") current, err := s.lockConnector(ctx, tx, c.String("id"), user, admin) if err != nil { return nil, err @@ -145,7 +149,10 @@ func (s *Service) testConnection(ctx context.Context, c, cred resource.Object, u } if cred != nil { fresh, err := resource.DecodeObject(sqlc.New(tx).LockCredential(ctx, cred.String("id"))) - if err != nil || fresh.Int("revision") != cred.Int("revision") || credentialStatus(current, fresh) != "authorized" { + if err != nil { + return nil, err + } + if fresh.Int("revision") != cred.Int("revision") || credentialStatus(current, fresh) != "authorized" { return nil, resource.Conflict } } @@ -168,7 +175,10 @@ func (s *Service) testConnection(ctx context.Context, c, cred resource.Object, u if err != nil { break } - schema, _ := json.Marshal(tool.InputSchema) + schema, marshalErr := json.Marshal(tool.InputSchema) + if marshalErr != nil { + return nil, fmt.Errorf("编码工具 %s 的输入结构: %w", tool.Name, marshalErr) + } _, err = queries.UpsertTool(ctx, sqlc.UpsertToolParams{ ConnectorID: c.String("id"), CredentialID: cred.String("id"), Name: tool.Name, Description: tool.Description, InputSchema: schema, diff --git a/monkeyai/backend/internal/model/admin.go b/monkeyai/backend/internal/model/admin.go index 49ba0daf5..49049fadd 100644 --- a/monkeyai/backend/internal/model/admin.go +++ b/monkeyai/backend/internal/model/admin.go @@ -3,6 +3,7 @@ package model import ( "encoding/json" "errors" + "log/slog" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity" @@ -12,9 +13,15 @@ import ( func (s *Service) RegisterAdmin(router chi.Router) { router.Get("/models", func(w http.ResponseWriter, r *http.Request) { - models, err := s.List(r.Context(), r.URL.Query().Get("ownership_type")) + ownership := r.URL.Query().Get("ownership_type") + models, err := s.List(r.Context(), ownership) if err != nil { - modelError(w, http.StatusBadRequest, err.Error()) + if ownership != "" && ownership != "system" && ownership != "user" { + modelError(w, http.StatusBadRequest, err.Error()) + } else { + slog.ErrorContext(r.Context(), "读取模型列表失败", "error", err) + modelError(w, http.StatusInternalServerError, "读取模型列表失败") + } return } for i := range models { @@ -25,6 +32,7 @@ func (s *Service) RegisterAdmin(router chi.Router) { router.Get("/models/authorization-subjects", func(w http.ResponseWriter, r *http.Request) { subjects, err := s.Subjects(r.Context()) if err != nil { + slog.ErrorContext(r.Context(), "读取模型授权对象失败", "error", err) modelError(w, http.StatusInternalServerError, "读取授权对象失败") return } @@ -76,6 +84,8 @@ func (s *Service) updateModel(w http.ResponseWriter, r *http.Request) { status := http.StatusBadRequest if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "操作模型失败", "model_id", chi.URLParam(r, "modelID"), "error", err) } modelError(w, status, err.Error()) return @@ -96,6 +106,8 @@ func (s *Service) setModelEnabled(w http.ResponseWriter, r *http.Request) { status := http.StatusInternalServerError if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "操作模型失败", "model_id", chi.URLParam(r, "modelID"), "error", err) } modelError(w, status, err.Error()) return @@ -108,6 +120,8 @@ func (s *Service) deleteModel(w http.ResponseWriter, r *http.Request) { status := http.StatusInternalServerError if errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "操作模型失败", "model_id", chi.URLParam(r, "modelID"), "error", err) } modelError(w, status, err.Error()) return @@ -131,7 +145,9 @@ func decodeModelRequest(w http.ResponseWriter, r *http.Request, target any) erro func modelJSON(w http.ResponseWriter, status int, value any) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(value) + if err := json.NewEncoder(w).Encode(value); err != nil { + slog.Error("写入模型 HTTP 响应失败", "status", status, "error", err) + } } func modelError(w http.ResponseWriter, status int, message string) { diff --git a/monkeyai/backend/internal/model/agent.go b/monkeyai/backend/internal/model/agent.go index 226c6a07f..c9b1997d3 100644 --- a/monkeyai/backend/internal/model/agent.go +++ b/monkeyai/backend/internal/model/agent.go @@ -3,6 +3,7 @@ package model import ( "context" "errors" + "log/slog" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/httpapi" @@ -104,7 +105,7 @@ func (s *Service) RegisterAgent(router chi.Router) { "total_count": len(items), "page": page, "page_size": size, "model_gateway": map[string]string{"base_url": s.gatewayURL, "authentication": "api_key"}, }); err != nil { - userModelError(w, err) + slog.ErrorContext(r.Context(), "写入模型目录缓存响应失败", "user_id", user.ID, "error", err) } }) router.Get("/models/tags", func(w http.ResponseWriter, r *http.Request) { @@ -119,7 +120,7 @@ func (s *Service) RegisterAgent(router chi.Router) { rows = append(rows, resource.Object{"tags": item.Tags}) } if err := httpapi.CachedJSON(w, r, map[string]any{"tags": resource.CollectTags(rows)}); err != nil { - userModelError(w, err) + slog.ErrorContext(r.Context(), "写入模型标签缓存响应失败", "user_id", user.ID, "error", err) } }) router.Get("/models/{modelID}", func(w http.ResponseWriter, r *http.Request) { diff --git a/monkeyai/backend/internal/model/postgres.go b/monkeyai/backend/internal/model/postgres.go index d4a33d7b5..d2b224ac8 100644 --- a/monkeyai/backend/internal/model/postgres.go +++ b/monkeyai/backend/internal/model/postgres.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/database" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/model/sqlc" @@ -70,7 +71,11 @@ func (p *Postgres) Create(ctx context.Context, item Model) (Model, error) { if err != nil { return Model{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚创建模型事务失败", "model_id", item.ID, "error", err) + } + }() advanced, err := json.Marshal(item.AdvancedConfig) if err != nil { return Model{}, err @@ -130,7 +135,11 @@ func (p *Postgres) update(ctx context.Context, item Model, ownership string) (Mo if err != nil { return Model{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚更新模型事务失败", "model_id", item.ID, "error", err) + } + }() advanced, err := json.Marshal(item.AdvancedConfig) if err != nil { return Model{}, err @@ -220,7 +229,11 @@ func (p *Postgres) delete(ctx context.Context, id, ownership, userID string) err if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer func() { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && ctx.Err() == nil { + slog.ErrorContext(ctx, "回滚删除模型事务失败", "model_id", id, "error", err) + } + }() result, err := sqlc.New(tx).DeleteModel(ctx, sqlc.DeleteModelParams{ID: id, OwnershipType: ownership, OwnerUserID: userID}) if err != nil { return err diff --git a/monkeyai/backend/internal/model/service.go b/monkeyai/backend/internal/model/service.go index f5f4cea0a..17b872131 100644 --- a/monkeyai/backend/internal/model/service.go +++ b/monkeyai/backend/internal/model/service.go @@ -3,6 +3,7 @@ package model import ( "context" "errors" + "log/slog" "net/url" "slices" "strings" @@ -74,6 +75,10 @@ func (s *Service) validateImageCapability(item Model) error { if err != nil { return err } + return validateImageCapabilities(item, cap) +} + +func validateImageCapabilities(item Model, cap ImageCapabilities) error { for _, quality := range item.ImageConfig.Qualities { if !slices.Contains(cap.Qualities, quality) { return errors.New("画质档位不被上游模型支持") @@ -168,7 +173,11 @@ func (s *Service) AgentModels(ctx context.Context, userID string, isAdmin bool) cap := ImageCapabilities{} if s.imageCapabilities != nil { cap, err = s.imageCapabilities(item.Provider, item.ModelID) - if err != nil || s.validateImageCapability(item) != nil { + if err != nil { + slog.ErrorContext(ctx, "读取模型生图能力失败", "model_id", item.ID, "error", err) + continue + } + if validateImageCapabilities(item, cap) != nil { continue } } diff --git a/monkeyai/backend/internal/proxy/billing.go b/monkeyai/backend/internal/proxy/billing.go index 513c47abd..a840ea61e 100644 --- a/monkeyai/backend/internal/proxy/billing.go +++ b/monkeyai/backend/internal/proxy/billing.go @@ -3,6 +3,7 @@ package proxy import ( "context" "encoding/json" + "fmt" "net/http" "time" @@ -83,7 +84,11 @@ func prepareRequest(body []byte, path string, limit int64, stream bool) ([]byte, delete(data, "max_completion_tokens") delete(data, "max_tokens") } - data[field], _ = json.Marshal(limit) + value, err := json.Marshal(limit) + if err != nil { + return nil, fmt.Errorf("编码输出上限: %w", err) + } + data[field] = value } if stream && path == "/v1/chat/completions" { options := map[string]json.RawMessage{} @@ -93,7 +98,11 @@ func prepareRequest(body []byte, path string, limit int64, stream bool) ([]byte, } } options["include_usage"] = json.RawMessage("true") - data["stream_options"], _ = json.Marshal(options) + value, err := json.Marshal(options) + if err != nil { + return nil, fmt.Errorf("编码流式选项: %w", err) + } + data["stream_options"] = value } return json.Marshal(data) } diff --git a/monkeyai/backend/internal/proxy/proxy.go b/monkeyai/backend/internal/proxy/proxy.go index e52000acb..1c8780015 100644 --- a/monkeyai/backend/internal/proxy/proxy.go +++ b/monkeyai/backend/internal/proxy/proxy.go @@ -157,8 +157,8 @@ func (p *Proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) { } upstream, err := parseBaseURL(target.BaseURL) if err != nil { - p.logger.ErrorContext(r.Context(), "模型配置无效", "model_id", target.ModelID, "error", err) - p.errorHandler(w, r, err) + p.logger.ErrorContext(r.Context(), "模型配置无效", "model_id", target.ModelID, "operation", "parse_base_url", "error", fmt.Sprintf("%T", err)) + p.writeUpstreamFailure(w, r, target.ModelID, "") return } if target.UpstreamModel != "" && target.UpstreamModel != meta.Model { @@ -264,13 +264,23 @@ func (p *Proxy) rewrite(r *httputil.ProxyRequest) { } func (p *Proxy) errorHandler(w http.ResponseWriter, r *http.Request, err error) { + modelID, transactionID := "", "" if pc, ok := r.Context().Value(proxyContextKey{}).(*proxyContext); ok { + modelID, transactionID = pc.target.ModelID, pc.reservation.ID p.finish(r.Context(), pc, Call{Stream: pc.stream, Result: "failed", ErrorCode: "upstream_connection_failed"}) } - p.logger.ErrorContext(r.Context(), "模型上游请求失败", "path", r.URL.Path, "error", err) + if r.Context().Err() == nil || !errors.Is(err, context.Canceled) { + p.logger.ErrorContext(r.Context(), "模型上游请求失败", "model_id", modelID, "transaction_id", transactionID, "operation", "proxy_request", "path", r.URL.Path, "error", fmt.Sprintf("%T", err)) + } + p.writeUpstreamFailure(w, r, modelID, transactionID) +} + +func (p *Proxy) writeUpstreamFailure(w http.ResponseWriter, r *http.Request, modelID, transactionID string) { w.Header().Set("Content-Type", "text/plain; charset=utf-8") w.WriteHeader(http.StatusBadGateway) - _, _ = io.WriteString(w, upstreamFailureMessage) + if _, writeErr := io.WriteString(w, upstreamFailureMessage); writeErr != nil && r.Context().Err() == nil && !errors.Is(writeErr, io.ErrClosedPipe) { + p.logger.WarnContext(r.Context(), "写入代理错误响应失败", "model_id", modelID, "transaction_id", transactionID, "operation", "write_error_response", "error", fmt.Sprintf("%T", writeErr)) + } } func upstreamPath(requestPath string) (string, bool) { diff --git a/monkeyai/backend/internal/proxy/proxy_test.go b/monkeyai/backend/internal/proxy/proxy_test.go index 95d664d4c..8df1501d0 100644 --- a/monkeyai/backend/internal/proxy/proxy_test.go +++ b/monkeyai/backend/internal/proxy/proxy_test.go @@ -1,6 +1,7 @@ package proxy import ( + "bytes" "context" "errors" "io" @@ -248,3 +249,20 @@ func (r *usageRecorderStub) Record(_ context.Context, call Call) error { r.calls <- call return nil } + +func TestProxyErrorHandlerRedactsUpstreamError(t *testing.T) { + var logs bytes.Buffer + proxy := NewProxy(nil, slog.New(slog.NewTextHandler(&logs, nil))) + pc := &proxyContext{target: Target{ModelID: "model-1"}, reservation: Reservation{ID: "transaction-1"}} + req := httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + req = req.WithContext(context.WithValue(req.Context(), proxyContextKey{}, pc)) + recorder := httptest.NewRecorder() + proxy.errorHandler(recorder, req, errors.New("upstream-token-and-private-response")) + if recorder.Code != http.StatusBadGateway || recorder.Body.String() != upstreamFailureMessage { + t.Fatalf("代理错误响应不正确: %d %q", recorder.Code, recorder.Body.String()) + } + text := logs.String() + if !strings.Contains(text, "model-1") || !strings.Contains(text, "transaction-1") || !strings.Contains(text, "proxy_request") || strings.Contains(text, "upstream-token-and-private-response") { + t.Fatalf("代理日志缺少安全上下文或泄漏上游详情: %s", text) + } +} diff --git a/monkeyai/backend/internal/proxy/reconcile.go b/monkeyai/backend/internal/proxy/reconcile.go index 8c81d2c57..01fc28023 100644 --- a/monkeyai/backend/internal/proxy/reconcile.go +++ b/monkeyai/backend/internal/proxy/reconcile.go @@ -2,8 +2,10 @@ package proxy import ( "context" + "errors" "fmt" "io" + "log/slog" "net/http" "strings" ) @@ -63,7 +65,12 @@ func (r *ResponseReconciler) Reconcile(ctx context.Context, target Target, respo if err != nil { return ResponseReconciliation{}, err } - defer response.Body.Close() + defer func() { + if closeErr := response.Body.Close(); closeErr != nil && ctx.Err() == nil && !errors.Is(closeErr, context.Canceled) && !errors.Is(closeErr, io.ErrClosedPipe) { + // 自定义 Transport 的关闭错误可能包含凭据或响应内容。 + slog.WarnContext(ctx, "关闭上游对账响应失败", "model_id", target.ModelID, "operation", "close_reconciliation_response", "error_type", fmt.Sprintf("%T", closeErr)) + } + }() if response.StatusCode != http.StatusOK { switch response.StatusCode { diff --git a/monkeyai/backend/internal/proxy/reconcile_test.go b/monkeyai/backend/internal/proxy/reconcile_test.go index 99ad7e97f..81c35385c 100644 --- a/monkeyai/backend/internal/proxy/reconcile_test.go +++ b/monkeyai/backend/internal/proxy/reconcile_test.go @@ -1,9 +1,14 @@ package proxy import ( + "bytes" "context" + "errors" + "io" + "log/slog" "net/http" "net/http/httptest" + "strings" "testing" ) @@ -66,3 +71,41 @@ func TestResponseReconcilerKeepsUncertainResultsPending(t *testing.T) { }) } } + +type responseCloseTransport struct { + body io.ReadCloser +} + +func (t responseCloseTransport) RoundTrip(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: t.body}, nil +} + +func TestResponseReconciliationCloseLogIsSafe(t *testing.T) { + for _, tc := range []struct { + name string + closeErr error + logged bool + }{ + {name: "unexpected close error", closeErr: errors.New("private-upstream-key-and-body"), logged: true}, + {name: "closed pipe", closeErr: io.ErrClosedPipe}, + {name: "canceled", closeErr: context.Canceled}, + } { + t.Run(tc.name, func(t *testing.T) { + var logs bytes.Buffer + original := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(original) }) + body := closeErrorReader{Reader: strings.NewReader(`{"id":"resp_test","status":"completed","usage":{"input_tokens":1,"output_tokens":1}}`), err: tc.closeErr} + target := testTarget("https://example.invalid/v1") + target.Protocol = "openai_responses" + result, err := NewResponseReconciler().WithTransport(responseCloseTransport{body: body}).Reconcile(context.Background(), target, "resp_test") + if err != nil || result.State != ResponseReconciliationResolved { + t.Fatalf("响应关闭不能改变对账结果: %+v %v", result, err) + } + text := logs.String() + if strings.Contains(text, "private-upstream-key-and-body") || tc.logged && (!strings.Contains(text, "close_reconciliation_response") || !strings.Contains(text, "model-1")) || !tc.logged && text != "" { + t.Fatalf("关闭响应时的日志不安全或有误报: %s", text) + } + }) + } +} diff --git a/monkeyai/backend/internal/proxy/usage.go b/monkeyai/backend/internal/proxy/usage.go index f80ddf153..f0d827f98 100644 --- a/monkeyai/backend/internal/proxy/usage.go +++ b/monkeyai/backend/internal/proxy/usage.go @@ -5,6 +5,8 @@ import ( "bytes" "context" "encoding/json" + "errors" + "fmt" "io" "log/slog" "net/http" @@ -59,11 +61,12 @@ type usageCaptureContext struct { } type usageCapture struct { - logger *slog.Logger - src io.ReadCloser - ctx usageCaptureContext - reader *io.PipeReader - writer *io.PipeWriter + logger *slog.Logger + src io.ReadCloser + ctx usageCaptureContext + reader *io.PipeReader + writer *io.PipeWriter + copyFailed bool } var _ io.ReadCloser = (*usageCapture)(nil) @@ -164,7 +167,11 @@ func (p *Proxy) recordUsage(ctx context.Context, proxyCtx *proxyContext, result } func (c *usageCapture) handleShadow() { - defer c.reader.Close() + defer func() { + if err := c.reader.Close(); err != nil { + c.logPipeError("close_usage_reader", err) + } + }() var result usageResult if c.ctx.stream { result = c.handleStream() @@ -388,18 +395,46 @@ func (c *usageCapture) handleNonStream() usageResult { return result } +func (c *usageCapture) logPipeError(operation string, err error) { + if errors.Is(err, io.ErrClosedPipe) || errors.Is(err, context.Canceled) || (c.ctx.ctx != nil && c.ctx.ctx.Err() != nil) { + return + } + transactionID, modelID := "", "" + if c.ctx.proxyCtx != nil { + transactionID = c.ctx.proxyCtx.reservation.ID + modelID = c.ctx.proxyCtx.target.ModelID + } + c.logger.WarnContext(c.ctx.ctx, "转发模型用量副本失败", "transaction_id", transactionID, "model_id", modelID, "operation", operation, "path", c.ctx.path, "error", err) +} + func (c *usageCapture) Close() error { - _ = c.writer.CloseWithError(io.ErrUnexpectedEOF) - return c.src.Close() + if err := c.writer.CloseWithError(io.ErrUnexpectedEOF); err != nil && !errors.Is(err, io.ErrClosedPipe) { + c.logPipeError("close_usage_pipe", err) + } + if err := c.src.Close(); err != nil { + // 上游 Body 可由自定义 Transport 提供,Close 错误文本未必可安全记录。 + if !errors.Is(err, io.ErrClosedPipe) && !errors.Is(err, context.Canceled) { + c.logPipeError("close_upstream_response", fmt.Errorf("上游响应关闭错误类型 %T", err)) + } + return fmt.Errorf("关闭上游响应: %w", err) + } + return nil } func (c *usageCapture) Read(buffer []byte) (int, error) { n, err := c.src.Read(buffer) - if n > 0 { - _, _ = c.writer.Write(bytes.Clone(buffer[:n])) + if n > 0 && !c.copyFailed { + if _, writeErr := c.writer.Write(bytes.Clone(buffer[:n])); writeErr != nil { + c.copyFailed = true + if c.ctx.ctx.Err() == nil { + c.logPipeError("write_usage_pipe", writeErr) + } + } } if err != nil { - _ = c.writer.CloseWithError(err) + if closeErr := c.writer.CloseWithError(err); closeErr != nil && !errors.Is(closeErr, io.ErrClosedPipe) { + c.logPipeError("close_usage_pipe", closeErr) + } } return n, err } diff --git a/monkeyai/backend/internal/proxy/usage_test.go b/monkeyai/backend/internal/proxy/usage_test.go index eb5826f80..363a1777c 100644 --- a/monkeyai/backend/internal/proxy/usage_test.go +++ b/monkeyai/backend/internal/proxy/usage_test.go @@ -1,7 +1,9 @@ package proxy import ( + "bytes" "context" + "errors" "io" "log/slog" "strings" @@ -133,3 +135,136 @@ func TestUsageCaptureReadCopiesResponse(t *testing.T) { t.Fatalf("result = %+v", parsed) } } + +func TestUsageCaptureClosedPipeDoesNotInterruptResponse(t *testing.T) { + var logs bytes.Buffer + reader, writer := io.Pipe() + if err := reader.Close(); err != nil { + t.Fatal(err) + } + capture := &usageCapture{ + logger: slog.New(slog.NewTextHandler(&logs, nil)), + src: io.NopCloser(strings.NewReader("private-upstream-response")), + ctx: usageCaptureContext{ctx: context.Background(), path: "/v1/responses", proxyCtx: &proxyContext{ + target: Target{ModelID: "model-1"}, reservation: Reservation{ID: "transaction-1"}, + }}, + reader: reader, writer: writer, + } + body, err := io.ReadAll(capture) + if err != nil || string(body) != "private-upstream-response" { + t.Fatalf("转发响应失败: %q %v", body, err) + } + if err := capture.Close(); err != nil { + t.Fatal(err) + } + text := logs.String() + if text != "" { + t.Fatalf("正常关闭的管道不应记录错误: %s", text) + } +} + +func TestUsageCaptureNormalCloseDoesNotLogError(t *testing.T) { + var logs bytes.Buffer + reader, writer := io.Pipe() + capture := &usageCapture{ + logger: slog.New(slog.NewTextHandler(&logs, nil)), + src: io.NopCloser(strings.NewReader("")), + ctx: usageCaptureContext{ctx: context.Background(), path: "/v1/responses"}, + reader: reader, writer: writer, + } + if _, err := io.ReadAll(capture); err != nil { + t.Fatal(err) + } + if err := capture.Close(); err != nil { + t.Fatal(err) + } + if logs.Len() != 0 { + t.Fatalf("正常关闭被误判为错误: %s", logs.String()) + } + if err := reader.Close(); err != nil { + t.Fatal(err) + } +} + +func TestUsageCapturePipeErrorLogsDetailWithoutResponse(t *testing.T) { + var logs bytes.Buffer + reader, writer := io.Pipe() + if err := reader.CloseWithError(errors.New("用量解析器意外失败")); err != nil { + t.Fatal(err) + } + capture := &usageCapture{ + logger: slog.New(slog.NewTextHandler(&logs, nil)), + src: io.NopCloser(strings.NewReader("private-upstream-response")), + ctx: usageCaptureContext{ctx: context.Background(), path: "/v1/responses", proxyCtx: &proxyContext{ + target: Target{ModelID: "model-1"}, reservation: Reservation{ID: "transaction-1"}, + }}, + reader: reader, writer: writer, + } + body, err := io.ReadAll(capture) + if err != nil || string(body) != "private-upstream-response" { + t.Fatalf("响应转发失败: %q %v", body, err) + } + if err := capture.Close(); err != nil { + t.Fatal(err) + } + text := logs.String() + if strings.Count(text, "write_usage_pipe") != 1 || !strings.Contains(text, "用量解析器意外失败") || !strings.Contains(text, "transaction-1") || !strings.Contains(text, "model-1") || strings.Contains(text, "private-upstream-response") { + t.Fatalf("管道错误日志缺少详情或泄漏响应: %s", text) + } +} + +func TestUsageCaptureCanceledRequestDoesNotLogPipeError(t *testing.T) { + var logs bytes.Buffer + ctx, cancel := context.WithCancel(context.Background()) + reader, writer := io.Pipe() + if err := reader.CloseWithError(errors.New("管道意外关闭")); err != nil { + t.Fatal(err) + } + cancel() + capture := &usageCapture{ + logger: slog.New(slog.NewTextHandler(&logs, nil)), + src: io.NopCloser(strings.NewReader("upstream-response")), + ctx: usageCaptureContext{ctx: ctx, path: "/v1/responses"}, + reader: reader, writer: writer, + } + if _, err := io.ReadAll(capture); err != nil { + t.Fatal(err) + } + if err := capture.Close(); err != nil { + t.Fatal(err) + } + if logs.Len() != 0 { + t.Fatalf("请求取消不应记录管道错误: %s", logs.String()) + } +} + +type closeErrorReader struct { + io.Reader + err error +} + +func (r closeErrorReader) Close() error { return r.err } + +func TestUsageCaptureUpstreamCloseDoesNotLogPrivateError(t *testing.T) { + var logs bytes.Buffer + reader, writer := io.Pipe() + defer func() { + if err := reader.Close(); err != nil && !errors.Is(err, io.ErrClosedPipe) { + t.Errorf("关闭测试管道失败: %v", err) + } + }() + upstreamErr := errors.New("private-upstream-response") + capture := &usageCapture{ + logger: slog.New(slog.NewTextHandler(&logs, nil)), + src: closeErrorReader{Reader: strings.NewReader(""), err: upstreamErr}, + ctx: usageCaptureContext{ctx: context.Background(), path: "/v1/responses"}, + reader: reader, writer: writer, + } + if err := capture.Close(); !errors.Is(err, upstreamErr) { + t.Fatalf("上游关闭错误未上抛: %v", err) + } + text := logs.String() + if !strings.Contains(text, "close_upstream_response") || strings.Contains(text, "private-upstream-response") { + t.Fatalf("上游关闭日志不安全: %s", text) + } +} diff --git a/monkeyai/backend/internal/resource/admin.go b/monkeyai/backend/internal/resource/admin.go index d85f26bb2..e929f637c 100644 --- a/monkeyai/backend/internal/resource/admin.go +++ b/monkeyai/backend/internal/resource/admin.go @@ -69,7 +69,7 @@ func (s *Store) RegisterAdmin(r chi.Router) { Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer func() { rollback(r.Context(), tx, "tag_save", chi.URLParam(r, "id")) }() id := chi.URLParam(r, "id") var out Object if id == "" { @@ -99,7 +99,7 @@ func (s *Store) RegisterAdmin(r chi.Router) { Fail(w, err) return } - defer tx.Rollback(r.Context()) + defer func() { rollback(r.Context(), tx, "tag_delete", chi.URLParam(r, "id")) }() id := chi.URLParam(r, "id") o, err := DecodeObject(sqlc.New(tx).DeleteTag(r.Context(), id)) u, _ := identity.UserFromContext(r.Context()) diff --git a/monkeyai/backend/internal/resource/grants.go b/monkeyai/backend/internal/resource/grants.go index d1712bb81..6d63f39a8 100644 --- a/monkeyai/backend/internal/resource/grants.go +++ b/monkeyai/backend/internal/resource/grants.go @@ -41,7 +41,7 @@ func (s *Store) RegisterGrants(r chi.Router, resources map[string]*CRUD) { Fail(w, err) return } - defer tx.Rollback(ctx) + defer func() { rollback(ctx, tx, "update_grants", chi.URLParam(r, "id")) }() id := chi.URLParam(r, "id") o, err := DecodeObject(c.Def.Repository(tx).LockResource(ctx, id)) if err != nil { diff --git a/monkeyai/backend/internal/resource/sharing.go b/monkeyai/backend/internal/resource/sharing.go index b9a9c1909..86454b980 100644 --- a/monkeyai/backend/internal/resource/sharing.go +++ b/monkeyai/backend/internal/resource/sharing.go @@ -105,7 +105,7 @@ func (s *Store) Share(ctx context.Context, actor string, input ShareInput, revok if err != nil { return err } - defer tx.Rollback(ctx) + defer func() { rollback(ctx, tx, "share", input.Resources[0].ID) }() // 固定加锁顺序,批量授权和删除共享资源行锁。 for _, item := range input.Resources { if err := kinds[item.Type].LockOwned(ctx, tx, item.ID, actor); err != nil { diff --git a/monkeyai/backend/internal/resource/store.go b/monkeyai/backend/internal/resource/store.go index 04496d069..e30a536fa 100644 --- a/monkeyai/backend/internal/resource/store.go +++ b/monkeyai/backend/internal/resource/store.go @@ -36,14 +36,24 @@ func (o Object) Int(k string) int64 { v, _ := o[k].(float64); return int64(v func ID() string { return uuid.New().String() } func Hash(v any) string { - b, _ := json.Marshal(v) + b, err := json.Marshal(v) + if err != nil { + slog.Error("资源内容哈希编码失败", "type", fmt.Sprintf("%T", v), "error", err) + } h := sha256.Sum256(b) return "sha256:" + hex.EncodeToString(h[:]) } func Strings(v any) []string { result := []string{} - b, _ := json.Marshal(v) - _ = json.Unmarshal(b, &result) + b, err := json.Marshal(v) + if err != nil { + slog.Error("资源字符串列表编码失败", "type", fmt.Sprintf("%T", v), "error", err) + return result + } + if err := json.Unmarshal(b, &result); err != nil { + slog.Error("资源字符串列表解码失败", "error", err) + return []string{} + } if result == nil { return []string{} } @@ -67,7 +77,9 @@ var Conflict = &Error{Status: 412, Code: "revision_conflict", Message: "资源 func JSON(w http.ResponseWriter, status int, v any) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(v) + if err := json.NewEncoder(w).Encode(v); err != nil { + slog.Error("资源响应写入失败", "status", status, "error", err) + } } func Fail(w http.ResponseWriter, err error) { var e *Error @@ -140,7 +152,10 @@ func Grants(ctx context.Context, q Queryer, kind, id string) ([]Object, error) { return grants, nil } func SaveGrants(ctx context.Context, tx pgx.Tx, kind, id, actor string, raw any, personal bool) error { - b, _ := json.Marshal(raw) + b, err := json.Marshal(raw) + if err != nil { + return fmt.Errorf("编码 %s 资源 %s 的授权: %w", kind, id, err) + } var grants []struct { UserID string `json:"user_id"` GroupID string `json:"group_id"` @@ -283,6 +298,12 @@ func (c *CRUD) List(ctx context.Context, q Queryer) ([]Object, error) { } return out, nil } +func rollback(ctx context.Context, tx pgx.Tx, operation, id string) { + if err := tx.Rollback(ctx); err != nil && !errors.Is(err, pgx.ErrTxClosed) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + slog.ErrorContext(ctx, "回滚资源事务失败", "operation", operation, "resource_id", id, "error", err) + } +} + func (c *CRUD) Save(ctx context.Context, actor, id, match string, in Object) (Object, error) { return c.save(ctx, actor, id, match, in, false) } @@ -291,7 +312,7 @@ func (c *CRUD) save(ctx context.Context, actor, id, match string, in Object, per if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer func() { rollback(ctx, tx, "save", id) }() old := Object{} create := id == "" if create { @@ -390,7 +411,7 @@ func (c *CRUD) delete(ctx context.Context, actor, id, match string, personal boo if err != nil { return err } - defer tx.Rollback(ctx) + defer func() { rollback(ctx, tx, "delete", id) }() o, err := DecodeObject(c.Def.Repository(tx).LockResource(ctx, id)) if err != nil { return err @@ -567,7 +588,7 @@ func (c *CRUD) SetEnabled(ctx context.Context, actor, id, match string, enabled if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer func() { rollback(ctx, tx, "set_enabled", id) }() o, err := DecodeObject(c.Def.Repository(tx).LockResource(ctx, id)) if err != nil { return nil, err diff --git a/monkeyai/backend/internal/setting/admin.go b/monkeyai/backend/internal/setting/admin.go index 145c8540e..f86dd3d47 100644 --- a/monkeyai/backend/internal/setting/admin.go +++ b/monkeyai/backend/internal/setting/admin.go @@ -1,8 +1,10 @@ package setting import ( + "context" "encoding/json" "errors" + "log/slog" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity" @@ -13,10 +15,11 @@ func (s *Service) RegisterAdmin(router chi.Router) { router.Get("/settings", func(w http.ResponseWriter, r *http.Request) { records, err := s.AdminList(r.Context()) if err != nil { - settingError(w, http.StatusInternalServerError, "读取设置失败") + slog.ErrorContext(r.Context(), "读取设置列表失败", "error", err) + settingError(r.Context(), w, http.StatusInternalServerError, "读取设置失败") return } - settingJSON(w, http.StatusOK, map[string]any{"settings": records}) + settingJSON(r.Context(), w, http.StatusOK, map[string]any{"settings": records}) }) router.Get("/settings/{key}", func(w http.ResponseWriter, r *http.Request) { record, err := s.AdminGet(r.Context(), chi.URLParam(r, "key")) @@ -24,11 +27,13 @@ func (s *Service) RegisterAdmin(router chi.Router) { status := http.StatusInternalServerError if errors.Is(err, ErrUnknownKey) || errors.Is(err, ErrNotFound) { status = http.StatusNotFound + } else { + slog.ErrorContext(r.Context(), "读取设置失败", "key", chi.URLParam(r, "key"), "error", err) } - settingError(w, status, "设置不存在") + settingError(r.Context(), w, status, "设置不存在") return } - settingJSON(w, http.StatusOK, record) + settingJSON(r.Context(), w, http.StatusOK, record) }) router.Put("/settings/{key}", s.putSetting) router.Post("/settings/email/test", s.testEmail) @@ -40,29 +45,32 @@ func (s *Service) putSetting(w http.ResponseWriter, r *http.Request) { SchemaVersion int `json:"schema_version"` } if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 2<<20)).Decode(&input); err != nil { - settingError(w, http.StatusBadRequest, "请求格式无效") + settingError(r.Context(), w, http.StatusBadRequest, "请求格式无效") return } user, _ := identity.UserFromContext(r.Context()) record, err := s.Put(r.Context(), chi.URLParam(r, "key"), input.Value, input.SchemaVersion, user.ID) if err != nil { - settingError(w, http.StatusBadRequest, err.Error()) + settingError(r.Context(), w, http.StatusBadRequest, err.Error()) return } record.Value, err = redact(record.Key, record.Value) if err != nil { - settingError(w, http.StatusInternalServerError, "设置已保存,但响应脱敏失败") + slog.ErrorContext(r.Context(), "设置响应脱敏失败", "key", record.Key, "error", err) + settingError(r.Context(), w, http.StatusInternalServerError, "设置已保存,但响应脱敏失败") return } - settingJSON(w, http.StatusOK, record) + settingJSON(r.Context(), w, http.StatusOK, record) } -func settingJSON(w http.ResponseWriter, status int, value any) { +func settingJSON(ctx context.Context, w http.ResponseWriter, status int, value any) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(status) - _ = json.NewEncoder(w).Encode(value) + if err := json.NewEncoder(w).Encode(value); err != nil { + slog.WarnContext(ctx, "写入设置响应失败", "error", err) + } } -func settingError(w http.ResponseWriter, status int, message string) { - settingJSON(w, status, map[string]any{"error": map[string]string{"code": "setting_error", "message": message}}) +func settingError(ctx context.Context, w http.ResponseWriter, status int, message string) { + settingJSON(ctx, w, status, map[string]any{"error": map[string]string{"code": "setting_error", "message": message}}) } diff --git a/monkeyai/backend/internal/setting/agent.go b/monkeyai/backend/internal/setting/agent.go index 1bf91de57..d0ad7f256 100644 --- a/monkeyai/backend/internal/setting/agent.go +++ b/monkeyai/backend/internal/setting/agent.go @@ -1,6 +1,7 @@ package setting import ( + "log/slog" "net/http" "github.com/chaitin/MonkeyCode/monkeyai/backend/internal/httpapi" @@ -14,7 +15,8 @@ func (s *Service) RegisterAgent(router chi.Router) { err = httpapi.CachedJSON(w, r, map[string]any{"settings": config.Settings}) } if err != nil { - settingError(w, http.StatusInternalServerError, "读取设置失败") + slog.ErrorContext(r.Context(), "读取 Agent 设置失败", "error", err) + settingError(r.Context(), w, http.StatusInternalServerError, "读取设置失败") } }) } diff --git a/monkeyai/backend/internal/setting/email.go b/monkeyai/backend/internal/setting/email.go index 6bd5b7d87..be4df6533 100644 --- a/monkeyai/backend/internal/setting/email.go +++ b/monkeyai/backend/internal/setting/email.go @@ -7,11 +7,13 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "mime" "net" "net/http" "net/mail" "net/smtp" + "net/textproto" "strconv" "strings" "time" @@ -68,8 +70,16 @@ func sendEmail(ctx context.Context, cfg emailConfig, to, subject, body string) e if err != nil { return err } - defer conn.Close() - stop := context.AfterFunc(ctx, func() { _ = conn.Close() }) + defer func() { + if err := conn.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + slog.WarnContext(ctx, "关闭 SMTP 连接失败", "failure", smtpFailure(err)) + } + }() + stop := context.AfterFunc(ctx, func() { + if err := conn.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + slog.WarnContext(ctx, "取消 SMTP 连接时关闭失败", "failure", smtpFailure(err)) + } + }) defer stop() deadline, _ := ctx.Deadline() if err := conn.SetDeadline(deadline); err != nil { @@ -88,7 +98,11 @@ func sendEmail(ctx context.Context, cfg emailConfig, to, subject, body string) e if err != nil { return err } - defer client.Close() + defer func() { + if err := client.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + slog.WarnContext(ctx, "关闭 SMTP 客户端失败", "failure", smtpFailure(err)) + } + }() if cfg.Encryption == "starttls" { if err := client.StartTLS(tlsConfig); err != nil { return err @@ -124,7 +138,9 @@ func sendEmail(ctx context.Context, cfg emailConfig, to, subject, body string) e return err } // DATA 已确认接收后,QUIT 失败不应误报发送失败。 - _ = client.Quit() + if err := client.Quit(); err != nil && !errors.Is(err, net.ErrClosed) { + slog.WarnContext(ctx, "SMTP QUIT 失败", "failure", smtpFailure(err)) + } return nil } @@ -133,18 +149,38 @@ func (s *Service) testEmail(w http.ResponseWriter, r *http.Request) { Recipient string `json:"recipient"` } if json.NewDecoder(http.MaxBytesReader(w, r.Body, 4096)).Decode(&input) != nil { - settingError(w, 400, "请求格式无效") + settingError(r.Context(), w, 400, "请求格式无效") return } input.Recipient = strings.TrimSpace(input.Recipient) address, err := mail.ParseAddress(input.Recipient) if err != nil || address.Address != input.Recipient { - settingError(w, 400, "收件邮箱无效") + settingError(r.Context(), w, 400, "收件邮箱无效") return } if err := s.Send(r.Context(), input.Recipient, "MonkeyAI 测试邮件", "这是一封 MonkeyAI 测试邮件,SMTP 发件配置已生效。"); err != nil { - settingError(w, http.StatusBadGateway, "邮件发送失败,请检查已保存的 SMTP 配置和服务连接") + slog.ErrorContext(r.Context(), "测试邮件发送失败", "failure", smtpFailure(err)) + settingError(r.Context(), w, http.StatusBadGateway, "邮件发送失败,请检查已保存的 SMTP 配置和服务连接") return } - settingJSON(w, http.StatusOK, map[string]bool{"sent": true}) + settingJSON(r.Context(), w, http.StatusOK, map[string]bool{"sent": true}) +} + +// SMTP 响应和网络错误可能回显收件人或认证信息,只记录协议状态及错误类别。 +func smtpFailure(err error) []any { + var response *textproto.Error + if errors.As(err, &response) { + return []any{"reason", "smtp_rejected", "smtp_status", response.Code} + } + if errors.Is(err, ErrNotFound) { + return []any{"error", ErrNotFound} + } + if errors.Is(err, context.DeadlineExceeded) { + return []any{"reason", "timeout"} + } + var network net.Error + if errors.As(err, &network) { + return []any{"reason", "network_error", "timeout", network.Timeout()} + } + return []any{"reason", "email_error", "error_type", fmt.Sprintf("%T", err)} } diff --git a/monkeyai/backend/internal/setting/email_test.go b/monkeyai/backend/internal/setting/email_test.go index 00e9603e6..7245f30eb 100644 --- a/monkeyai/backend/internal/setting/email_test.go +++ b/monkeyai/backend/internal/setting/email_test.go @@ -148,3 +148,10 @@ func TestDefaultAuthenticationAndEmailSecret(t *testing.T) { } } } + +func TestSMTPFailureExcludesResponse(t *testing.T) { + logged := fmt.Sprint(smtpFailure(&textproto.Error{Code: 535, Msg: "secret@example.com private-token"})) + if strings.Contains(logged, "private-token") || strings.Contains(logged, "secret@example.com") || !strings.Contains(logged, "535") { + t.Fatalf("SMTP 日志分类不安全: %s", logged) + } +} diff --git a/monkeyai/backend/internal/setting/service.go b/monkeyai/backend/internal/setting/service.go index 0c12a346e..f9e3ac886 100644 --- a/monkeyai/backend/internal/setting/service.go +++ b/monkeyai/backend/internal/setting/service.go @@ -86,11 +86,22 @@ func (s *Service) Put(ctx context.Context, key string, value json.RawMessage, sc } id, ok := connection["id"] if !ok || string(id) == `""` { - connection["id"], _ = json.Marshal(resource.ID()) + encoded, err := json.Marshal(resource.ID()) + if err != nil { + return Record{}, fmt.Errorf("生成 OAuth 连接标识: %w", err) + } + connection["id"] = encoded } } - object["oauth_connections"], _ = json.Marshal(connections) - value, _ = json.Marshal(object) + encoded, err := json.Marshal(connections) + if err != nil { + return Record{}, fmt.Errorf("编码 OAuth 连接配置: %w", err) + } + object["oauth_connections"] = encoded + value, err = json.Marshal(object) + if err != nil { + return Record{}, fmt.Errorf("编码认证设置: %w", err) + } } } value, err := s.mergeSecrets(ctx, key, value) @@ -153,14 +164,19 @@ func (s *Service) AgentConfig(ctx context.Context) (Config, error) { } if record.Key == "billing" { var all map[string]json.RawMessage - _ = json.Unmarshal(value, &all) + if err := json.Unmarshal(value, &all); err != nil { + return Config{}, fmt.Errorf("读取计费设置: %w", err) + } safe := map[string]json.RawMessage{} for _, key := range []string{"input_credits_per_million_tokens", "cached_input_credits_per_million_tokens", "output_credits_per_million_tokens", "charging_mode", "enabled", "quota_refresh_cycle"} { if v, ok := all[key]; ok { safe[key] = v } } - value, _ = json.Marshal(safe) + value, err = json.Marshal(safe) + if err != nil { + return Config{}, fmt.Errorf("编码计费设置: %w", err) + } } config.Settings[record.Key] = value if record.UpdatedAt.After(config.UpdatedAt) { @@ -202,8 +218,11 @@ func (s *Service) mergeSecrets(ctx context.Context, key string, value json.RawMe return nil, err } var next, previous map[string]any - if json.Unmarshal(value, &next) != nil || json.Unmarshal(existing.Value, &previous) != nil { - return value, nil + if err := json.Unmarshal(value, &next); err != nil { + return nil, errors.New("value 必须是 JSON 对象") + } + if err := json.Unmarshal(existing.Value, &previous); err != nil { + return nil, fmt.Errorf("解析已有 %s 设置: %w", key, err) } switch key { case "authentication": @@ -268,7 +287,10 @@ func validate(key string, value map[string]json.RawMessage) error { } } case "email": - data, _ := json.Marshal(value) + data, err := json.Marshal(value) + if err != nil { + return errors.New("邮件配置格式无效") + } var config emailConfig if json.Unmarshal(data, &config) != nil { return errors.New("邮件配置格式无效") @@ -289,7 +311,9 @@ func validate(key string, value map[string]json.RawMessage) error { func rawString(value json.RawMessage) string { var result string - _ = json.Unmarshal(value, &result) + if err := json.Unmarshal(value, &result); err != nil { + return "" + } return result } diff --git a/monkeyai/backend/internal/setting/service_test.go b/monkeyai/backend/internal/setting/service_test.go index 9d367a5b7..090a00300 100644 --- a/monkeyai/backend/internal/setting/service_test.go +++ b/monkeyai/backend/internal/setting/service_test.go @@ -56,6 +56,21 @@ func TestAgentConfigRedactsSecrets(t *testing.T) { } } +func TestPutRejectsBrokenStoredSecrets(t *testing.T) { + broken := json.RawMessage(`{"smtp_password":"secret"`) + store := &memoryStore{records: map[string]Record{ + "email": {Key: "email", Value: broken}, + }} + service := NewService(store) + _, err := service.Put(t.Context(), "email", json.RawMessage(`{"smtp_host":"smtp.example.com","smtp_port":587,"smtp_encryption":"starttls","sender_email":"sender@example.com"}`), 1, "admin") + if err == nil { + t.Fatal("已保存的密钥损坏时不应覆盖现有配置") + } + if string(store.records["email"].Value) != string(broken) { + t.Fatal("损坏的已有配置被覆盖") + } +} + func TestPutPreservesRedactedSecret(t *testing.T) { store := &memoryStore{records: map[string]Record{ "authentication": {Key: "authentication", Value: json.RawMessage(`{"oauth_connections":[{"id":"github","provider":"github","name":"GitHub","client_id":"client","client_secret":"secret","enabled":true}]}`)}, diff --git a/monkeyai/backend/internal/skill/service.go b/monkeyai/backend/internal/skill/service.go index 545d9e148..3138aaa6e 100644 --- a/monkeyai/backend/internal/skill/service.go +++ b/monkeyai/backend/internal/skill/service.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "io" + "log/slog" "net/http" "strconv" @@ -34,7 +35,11 @@ func (s *Service) read(ctx context.Context, object resource.Object) (Package, er if err != nil { return Package{}, err } - defer r.Close() + defer func() { + if err := r.Close(); err != nil { + slog.ErrorContext(ctx, "关闭技能包读取流失败", "skill_id", object.String("id"), "error", err) + } + }() b, err := io.ReadAll(io.LimitReader(r, MaxPackage+1)) if err != nil { return Package{}, err @@ -155,13 +160,21 @@ func (s *Service) upload(w http.ResponseWriter, r *http.Request, personal bool) resource.Fail(w, resource.Invalid("技能包上传无效或超限")) return } - defer r.MultipartForm.RemoveAll() + defer func() { + if err := r.MultipartForm.RemoveAll(); err != nil { + slog.ErrorContext(r.Context(), "清理技能包上传临时文件失败", "error", err) + } + }() f, _, err := r.FormFile("package") if err != nil { resource.Fail(w, resource.Invalid("缺少 package")) return } - defer f.Close() + defer func() { + if err := f.Close(); err != nil { + slog.ErrorContext(r.Context(), "关闭技能包上传文件失败", "error", err) + } + }() b, err := io.ReadAll(io.LimitReader(f, MaxPackage+1)) if err != nil { resource.Fail(w, err) @@ -223,5 +236,7 @@ func (s *Service) ServePackage(w http.ResponseWriter, r *http.Request, o resourc w.Header().Set("ETag", `"`+o.String("package_sha256")+`"`) w.Header().Set("Cache-Control", "private, no-store") w.Header().Set("X-Content-SHA256", o.String("package_sha256")) - _, _ = w.Write(p.Bytes) + if _, err := w.Write(p.Bytes); err != nil { + slog.ErrorContext(r.Context(), "技能包下载写入失败", "skill_id", o.String("id"), "error", err) + } } diff --git a/monkeyai/backend/internal/stats/history.go b/monkeyai/backend/internal/stats/history.go index 69873dec1..cb05f987b 100644 --- a/monkeyai/backend/internal/stats/history.go +++ b/monkeyai/backend/internal/stats/history.go @@ -53,7 +53,7 @@ func (s *Service) history(w http.ResponseWriter, r *http.Request) { resource.Fail(w, err) return } - defer tx.Rollback(context.WithoutCancel(ctx)) + defer func() { rollbackStats(ctx, tx, "history") }() var total int64 total, err = sqlc.New(tx).CountHistory(ctx, sqlc.CountHistoryParams{TitleQuery: title, UserQuery: user, FromTime: from, UntilTime: until}) if err != nil { diff --git a/monkeyai/backend/internal/stats/service.go b/monkeyai/backend/internal/stats/service.go index 22b12c88f..047bb8ec1 100644 --- a/monkeyai/backend/internal/stats/service.go +++ b/monkeyai/backend/internal/stats/service.go @@ -2,6 +2,8 @@ package stats import ( "context" + "errors" + "log/slog" "net/http" "time" @@ -62,6 +64,12 @@ func (s *Service) RegisterAdmin(r chi.Router) { r.Get("/statistics/history", s.history) } +func rollbackStats(ctx context.Context, tx pgx.Tx, operation string) { + if err := tx.Rollback(context.WithoutCancel(ctx)); err != nil && !errors.Is(err, pgx.ErrTxClosed) { + slog.ErrorContext(ctx, "回滚统计查询事务失败", "operation", operation, "error", err) + } +} + func (s *Service) read(live bool, query func(context.Context, pgx.Tx, window) (resource.Object, error)) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { period, err := period(r, s.now().UTC(), live) @@ -76,7 +84,7 @@ func (s *Service) read(live bool, query func(context.Context, pgx.Tx, window) (r resource.Fail(w, err) return } - defer tx.Rollback(context.WithoutCancel(ctx)) + defer func() { rollbackStats(ctx, tx, "read") }() out, err := query(ctx, tx, period) if err != nil { resource.Fail(w, err)