Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 13 additions & 5 deletions monkeyai/backend/internal/agentconfig/resources.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@ package agentconfig

import (
"context"
"errors"
"fmt"
"log/slog"
"net/http"
"slices"
"strings"
Expand Down Expand Up @@ -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 {
Expand All @@ -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))
Expand All @@ -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)})
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions monkeyai/backend/internal/apikey/admin.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package apikey

import (
"errors"
"log/slog"
"net/http"

"github.com/go-chi/chi/v5"
Expand All @@ -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
}
Expand All @@ -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
Expand Down
10 changes: 9 additions & 1 deletion monkeyai/backend/internal/apikey/agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package apikey
import (
"encoding/json"
"errors"
"log/slog"
"net/http"

"github.com/chaitin/MonkeyCode/monkeyai/backend/internal/identity"
Expand All @@ -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
}
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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) {
Expand Down
5 changes: 4 additions & 1 deletion monkeyai/backend/internal/apikey/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"log/slog"

"github.com/chaitin/MonkeyCode/monkeyai/backend/internal/apikey/sqlc"

Expand Down Expand Up @@ -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
}
9 changes: 8 additions & 1 deletion monkeyai/backend/internal/apikey/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"encoding/hex"
"errors"
"fmt"
"log/slog"
"slices"
"strings"
"time"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
38 changes: 32 additions & 6 deletions monkeyai/backend/internal/apikey/service_test.go
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
}
Expand Down Expand Up @@ -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())
}
}
14 changes: 10 additions & 4 deletions monkeyai/backend/internal/audit/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand All @@ -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}
}
Expand All @@ -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)
Expand Down
19 changes: 13 additions & 6 deletions monkeyai/backend/internal/audit/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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 {
Expand All @@ -144,19 +146,24 @@ 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 {
item["request_params"] = sanitize(item["request_params"], 0)
}
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)
}
}
Loading
Loading