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
5 changes: 3 additions & 2 deletions monkeyai/backend/internal/app/resources_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ import (
"net/url"
"os"
"path/filepath"
"slices"
"strings"
"testing"

Expand Down Expand Up @@ -1320,8 +1321,8 @@ DEFERRABLE INITIALLY DEFERRED FOR EACH ROW EXECUTE FUNCTION reject_test_tag();`)
})

// 在已有业务数据的测试库验证完整回滚,再重新初始化。
for i := len(migrations) - 1; i >= 0; i-- {
path := strings.TrimSuffix(migrations[i], ".up.sql") + ".down.sql"
for _, migration := range slices.Backward(migrations) {
path := strings.TrimSuffix(migration, ".up.sql") + ".down.sql"
down, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
Expand Down
2 changes: 1 addition & 1 deletion monkeyai/backend/internal/billing/agent_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -219,7 +219,7 @@ func TestAgentBillingContract(t *testing.T) {
case map[string]any:
if ref, ok := node["$ref"].(string); ok && strings.HasPrefix(ref, "#/") {
var target any = document
for _, part := range strings.Split(strings.TrimPrefix(ref, "#/"), "/") {
for part := range strings.SplitSeq(strings.TrimPrefix(ref, "#/"), "/") {
parent, ok := target.(map[string]any)
if !ok || parent[part] == nil {
t.Fatalf("无效契约引用 %s", ref)
Expand Down
6 changes: 3 additions & 3 deletions monkeyai/backend/internal/endpoint/connection_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ func TestQueueBoundsAndSnapshots(t *testing.T) {
s := NewService(nil, nil, slog.New(slog.NewTextHandler(io.Discard, nil)), "")
c := queueConnection(s, nil, machineA)
defer c.cancel()
for i := 0; i < 64; i++ {
for range 64 {
if !c.enqueue([]byte("x")) {
t.Fatal("提前拒绝入队")
}
Expand Down Expand Up @@ -62,8 +62,8 @@ func TestConnectionFencingAndRouting(t *testing.T) {
defer old.cancel()
u.connections[machineA] = a
u.connections[machineB] = b
u.endpoints[machineA] = Endpoint{View: View{MachineID: machineA}}
u.endpoints[machineB] = Endpoint{View: View{MachineID: machineB}}
u.endpoints[machineA] = Endpoint{MachineID: machineA}
u.endpoints[machineB] = Endpoint{MachineID: machineB}
m := Message{Type: "event", ID: messageID, Target: machineB, Method: "agent.example", Payload: json.RawMessage(`{}`)}
s.route(old, m, 100)
if !old.dead.Load() || len(b.business) != 0 {
Expand Down
2 changes: 1 addition & 1 deletion monkeyai/backend/internal/endpoint/contract_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ func TestEndpointContract(t *testing.T) {
case map[string]any:
if ref, ok := v["$ref"].(string); ok && strings.HasPrefix(ref, "#/") {
var target any = document
for _, key := range strings.Split(strings.TrimPrefix(ref, "#/"), "/") {
for key := range strings.SplitSeq(strings.TrimPrefix(ref, "#/"), "/") {
node, ok := target.(map[string]any)
if !ok {
t.Fatalf("无效引用 %s", ref)
Expand Down
3 changes: 1 addition & 2 deletions monkeyai/backend/internal/endpoint/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -308,8 +308,7 @@ func (s *Service) connect(w http.ResponseWriter, r *http.Request) {
h, err := hello(data)
if err != nil {
name := "invalid_message"
var f fault
if errors.As(err, &f) {
if f, ok := errors.AsType[fault](err); ok {
name = f.code
}
c.sendError(name, "")
Expand Down
2 changes: 1 addition & 1 deletion monkeyai/backend/internal/endpoint/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ func fromRow(row sqlc.Endpoint) Endpoint {
if row.Alias != nil {
name = *row.Alias
}
return Endpoint{View: View{MachineID: row.MachineID, Profile: Profile{row.DeviceName, row.Platform, row.OsVersion, row.Arch, row.ClientVersion}, Alias: row.Alias, DisplayName: name, ProtocolVersion: row.ProtocolVersion, LastSeenAt: millis(row.LastSeenAt)}, Status: row.Status, CreatedAt: row.CreatedAt.UnixMilli(), UpdatedAt: row.UpdatedAt.UnixMilli(), RevokedAt: millis(row.RevokedAt)}
return Endpoint{MachineID: row.MachineID, Profile: Profile{row.DeviceName, row.Platform, row.OsVersion, row.Arch, row.ClientVersion}, Alias: row.Alias, DisplayName: name, ProtocolVersion: row.ProtocolVersion, LastSeenAt: millis(row.LastSeenAt), Status: row.Status, CreatedAt: row.CreatedAt.UnixMilli(), UpdatedAt: row.UpdatedAt.UnixMilli(), RevokedAt: millis(row.RevokedAt)}
}
func missing(err error) error {
if errors.Is(err, pgx.ErrNoRows) {
Expand Down
2 changes: 1 addition & 1 deletion monkeyai/backend/internal/endpoint/proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ func TestNginxUpgrade(t *testing.T) {
}
address := strings.TrimSpace(string(data))
var conn *websocket.Conn
for i := 0; i < 30; i++ {
for range 30 {
conn, _, err = websocket.Dial(ctx, "ws://"+address+"/api/v1/endpoints/connect", &websocket.DialOptions{HTTPHeader: http.Header{"Authorization": []string{"Bearer owner"}}})
if err == nil {
break
Expand Down
3 changes: 1 addition & 2 deletions monkeyai/backend/internal/endpoint/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -275,8 +275,7 @@ func (s *Service) check(ctx context.Context, credential Credential) (int, error)
return 0, nil
}
func status(err error) (int, string) {
var f fault
if errors.As(err, &f) {
if f, ok := errors.AsType[fault](err); ok {
switch f.code {
case "unauthorized":
return 401, f.code
Expand Down
3 changes: 1 addition & 2 deletions monkeyai/backend/internal/identity/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -334,8 +334,7 @@ func (s *Service) clientLoginURL(requestID, errorCode string) string {
}

func logUpstreamFailure(ctx context.Context, operation, connectionID string, err error) {
var response *upstreamHTTPError
if errors.As(err, &response) {
if response, ok := errors.AsType[*upstreamHTTPError](err); ok {
slog.ErrorContext(ctx, "上游 OAuth 操作失败", "operation", operation, "connection_id", connectionID, "reason", "http_status", "status", response.status)
return
}
Expand Down
9 changes: 3 additions & 6 deletions monkeyai/backend/internal/imagegen/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -312,9 +312,7 @@ func (s *Service) submit(ctx context.Context, target proxy.Target, operation, re
return imageproxy.Task{}, err
}
held = false
s.active.Add(1)
go func() {
defer s.active.Done()
s.active.Go(func() {
defer func() { <-s.slots }()
callCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Minute)
defer cancel()
Expand All @@ -327,7 +325,7 @@ func (s *Service) submit(ctx context.Context, target proxy.Target, operation, re
result, err = adapter.Editor.Edit(callCtx, target, request)
}
s.handleResult(callCtx, job, result, err)
}()
})
return imageproxy.Task{ID: job.ID, UserID: job.UserID, Operation: operation, Status: "pending",
Usage: &imageproxy.Usage{RequestedImages: imageCount}}, nil
}
Expand All @@ -346,8 +344,7 @@ 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) {
if _, ok := errors.AsType[*url.Error](err); ok {
return "url_error"
}
return fmt.Sprintf("%T", err)
Expand Down
21 changes: 7 additions & 14 deletions monkeyai/backend/internal/mcp/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,7 @@ type oauthConfig struct {
func oauthSettings(c resource.Object) oauthConfig {
b, err := json.Marshal(c["oauth_config"])
if err != nil {
var unsupported *json.UnsupportedTypeError
if errors.As(err, &unsupported) {
if unsupported, ok := errors.AsType[*json.UnsupportedTypeError](err); ok {
slog.Error("编码 OAuth 配置失败", "connector_id", c.String("id"), "operation", "encode_config", "error", unsupported)
} else {
// 自定义 JSON 错误可能回显配置中的客户端密钥。
Expand Down Expand Up @@ -374,19 +373,16 @@ 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, ok := errors.AsType[tokenExchangeError](err); ok {
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) {
if failure, ok := errors.AsType[*resource.Error](err); ok {
return []any{"reason", failure.Code, "status", failure.Status}
}
var upstream remoteStatus
if errors.As(err, &upstream) {
if upstream, ok := errors.AsType[remoteStatus](err); ok {
return []any{"reason", "upstream_http_error", "upstream_status", int(upstream)}
}
if errors.Is(err, context.DeadlineExceeded) {
Expand All @@ -395,16 +391,13 @@ func safeMCPFailure(err error) []any {
if errors.Is(err, context.Canceled) {
return []any{"reason", "canceled"}
}
var network net.Error
if errors.As(err, &network) {
if network, ok := errors.AsType[net.Error](err); ok {
return []any{"reason", "network_error", "timeout", network.Timeout()}
}
var urlErr *url.Error
if errors.As(err, &urlErr) {
if _, ok := errors.AsType[*url.Error](err); ok {
return []any{"reason", "transport_error"}
}
var databaseErr *pgconn.PgError
if errors.As(err, &databaseErr) {
if databaseErr, ok := errors.AsType[*pgconn.PgError](err); ok {
return []any{"reason", "database_error", "sqlstate", databaseErr.Code}
}
if errors.Is(err, pgx.ErrNoRows) {
Expand Down
3 changes: 1 addition & 2 deletions monkeyai/backend/internal/mcp/protocol.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,7 @@ func readRequest(w http.ResponseWriter, r *http.Request) (request, bool) {
body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, 1<<20))
if err != nil {
status := http.StatusBadRequest
var limit *http.MaxBytesError
if errors.As(err, &limit) {
if _, ok := errors.AsType[*http.MaxBytesError](err); ok {
status = http.StatusRequestEntityTooLarge
}
rpcReply(w, status, nil, nil, &rpcError{Code: -32700, Message: "请求超限或不完整"})
Expand Down
3 changes: 1 addition & 2 deletions monkeyai/backend/internal/mcp/tools.go
Original file line number Diff line number Diff line change
Expand Up @@ -160,8 +160,7 @@ func (s *Service) testConnection(ctx context.Context, c, cred resource.Object, u
status, message := "connected", ""
if discoveryErr != nil {
status, message = "error", "MCP 连接或工具发现失败,请检查地址和网络"
var upstreamStatus remoteStatus
if errors.As(discoveryErr, &upstreamStatus) {
if upstreamStatus, ok := errors.AsType[remoteStatus](discoveryErr); ok {
if upstreamStatus == http.StatusUnauthorized {
message = "上游认证失败,请更新 Header 或重新授权"
}
Expand Down
7 changes: 3 additions & 4 deletions monkeyai/backend/internal/model/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"errors"
"fmt"
"log/slog"
"slices"

"github.com/chaitin/MonkeyCode/monkeyai/backend/internal/database"
"github.com/chaitin/MonkeyCode/monkeyai/backend/internal/model/sqlc"
Expand Down Expand Up @@ -429,10 +430,8 @@ func normalizeAuthorization(value Authorization) Authorization {
if value.AllUsers {
return Authorization{AllUsers: true, UserIDs: []string{}, GroupIDs: []string{}}
}
for _, id := range value.GroupIDs {
if id == rootgroup.ID {
return Authorization{AllUsers: true, UserIDs: []string{}, GroupIDs: []string{}}
}
if slices.Contains(value.GroupIDs, rootgroup.ID) {
return Authorization{AllUsers: true, UserIDs: []string{}, GroupIDs: []string{}}
}
value.UserIDs = unique(value.UserIDs)
value.GroupIDs = unique(value.GroupIDs)
Expand Down
13 changes: 6 additions & 7 deletions monkeyai/backend/internal/setting/email.go
Original file line number Diff line number Diff line change
Expand Up @@ -124,14 +124,15 @@ func sendEmail(ctx context.Context, cfg emailConfig, to, subject, body string) e
return err
}
from := (&mail.Address{Name: cfg.SenderName, Address: cfg.SenderEmail}).String()
message := fmt.Sprintf("From: %s\r\nTo: %s\r\nSubject: %s\r\nDate: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: base64\r\n\r\n", from, recipient.String(), mime.QEncoding.Encode("UTF-8", subject), time.Now().Format(time.RFC1123Z))
var message strings.Builder
message.WriteString(fmt.Sprintf("From: %s\r\nTo: %s\r\nSubject: %s\r\nDate: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: base64\r\n\r\n", from, recipient.String(), mime.QEncoding.Encode("UTF-8", subject), time.Now().Format(time.RFC1123Z)))
encoded := base64.StdEncoding.EncodeToString([]byte(body))
for len(encoded) > 0 {
n := min(76, len(encoded))
message += encoded[:n] + "\r\n"
message.WriteString(encoded[:n] + "\r\n")
encoded = encoded[n:]
}
if _, err := writer.Write([]byte(message)); err != nil {
if _, err := writer.Write([]byte(message.String())); err != nil {
return err
}
if err := writer.Close(); err != nil {
Expand Down Expand Up @@ -168,8 +169,7 @@ func (s *Service) testEmail(w http.ResponseWriter, r *http.Request) {

// SMTP 响应和网络错误可能回显收件人或认证信息,只记录协议状态及错误类别。
func smtpFailure(err error) []any {
var response *textproto.Error
if errors.As(err, &response) {
if response, ok := errors.AsType[*textproto.Error](err); ok {
return []any{"reason", "smtp_rejected", "smtp_status", response.Code}
}
if errors.Is(err, ErrNotFound) {
Expand All @@ -178,8 +178,7 @@ func smtpFailure(err error) []any {
if errors.Is(err, context.DeadlineExceeded) {
return []any{"reason", "timeout"}
}
var network net.Error
if errors.As(err, &network) {
if network, ok := errors.AsType[net.Error](err); ok {
return []any{"reason", "network_error", "timeout", network.Timeout()}
}
return []any{"reason", "email_error", "error_type", fmt.Sprintf("%T", err)}
Expand Down
Loading