diff --git a/inventory/task.go b/inventory/task.go index ed2f291b..0f4d4ea3 100644 --- a/inventory/task.go +++ b/inventory/task.go @@ -3,6 +3,8 @@ package inventory import ( "context" "fmt" + "net/netip" + "strings" "time" "entgo.io/ent/dialect/sql" @@ -62,7 +64,8 @@ type ( Status []task.Status UserID int CorrelationID *uuid.UUID - // CreatorIP filters tasks created from a matching client IP (substring). + // CreatorIP filters tasks created from a matching client IP: CIDR + // notation ("10.0.0.0/8"), exact IP, or substring match. CreatorIP string // ExcludeHidden filters out tasks hidden by their owner. ExcludeHidden bool @@ -248,7 +251,30 @@ func (c *taskClient) List(ctx context.Context, args *ListTaskArgs) (*ListTaskRes } if args.CreatorIP != "" { - q.Where(task.CreatorIPContainsFold(args.CreatorIP)) + filter := strings.TrimSpace(args.CreatorIP) + if prefix, err := netip.ParsePrefix(filter); err == nil { + // CIDR filter: creator_ip is a string column and portable SQL has + // no INET functions, so resolve matching IPs first. The task table + // is bounded by periodic cleanup, keeping the IN list small. + ips, err := q.Clone().Select(task.FieldCreatorIP).Strings(ctx) + if err != nil { + return nil, fmt.Errorf("failed to resolve creator IPs: %w", err) + } + matched := make([]string, 0, len(ips)) + for _, ip := range ips { + if addr, aerr := netip.ParseAddr(ip); aerr == nil && prefix.Contains(addr) { + matched = append(matched, ip) + } + } + if len(matched) > c.maxSQlParam { + return nil, fmt.Errorf("IP range %q matches too many addresses, narrow the range", args.CreatorIP) + } + q.Where(task.CreatorIPIn(matched...)) + } else if addr, err := netip.ParseAddr(filter); err == nil { + q.Where(task.CreatorIPEQ(addr.String())) + } else { + q.Where(task.CreatorIPContainsFold(args.CreatorIP)) + } } if args.ExcludeHidden { diff --git a/inventory/task_ip_filter_test.go b/inventory/task_ip_filter_test.go new file mode 100644 index 00000000..a3f4cee3 --- /dev/null +++ b/inventory/task_ip_filter_test.go @@ -0,0 +1,60 @@ +package inventory + +import ( + "context" + "testing" + + "github.com/cloudreve/Cloudreve/v4/ent/enttest" + "github.com/cloudreve/Cloudreve/v4/inventory/types" + "github.com/cloudreve/Cloudreve/v4/pkg/conf" + "github.com/cloudreve/Cloudreve/v4/pkg/hashid" + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" +) + +func TestListTaskCreatorIPFilter(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + ctx := context.Background() + + mk := func(ip string) { + client.Task.Create(). + SetType("test"). + SetPublicState(&types.TaskPublicState{}). + SetCorrelationID(uuid.Must(uuid.NewV4())). + SetCreatorIP(ip). + SaveX(ctx) + } + mk("192.168.1.10") + mk("192.168.1.20") + mk("10.0.0.5") + mk("2001:db8::1") + + hasher, err := hashid.New("test") + require.NoError(t, err) + tc := NewTaskClient(client, conf.SQLite3DB, hasher) + + list := func(filter string) []string { + res, err := tc.List(ctx, &ListTaskArgs{ + PaginationArgs: &PaginationArgs{Page: 0, PageSize: 50}, + CreatorIP: filter, + }) + require.NoError(t, err) + ips := make([]string, 0, len(res.Tasks)) + for _, task := range res.Tasks { + ips = append(ips, task.CreatorIP) + } + return ips + } + + // CIDR prefix matches only in-range addresses. + require.ElementsMatch(t, []string{"192.168.1.10", "192.168.1.20"}, list("192.168.1.0/24")) + require.ElementsMatch(t, []string{"2001:db8::1"}, list("2001:db8::/64")) + + // Exact IP normalizes and matches. + require.ElementsMatch(t, []string{"10.0.0.5"}, list("10.0.0.5")) + + // Non-IP input falls back to substring matching. + require.ElementsMatch(t, []string{"192.168.1.10", "192.168.1.20"}, list("192.168")) + require.Empty(t, list("192.168.1.0/30")) +}