Merge pull request #138 from Dvorinka/merge/upstream-prs
Merge 7 upstream PRs + repair stale test suite + fork roadmappull/3579/head
commit
b347f036e8
@ -0,0 +1,83 @@
|
|||||||
|
package inventory
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent/enttest"
|
||||||
|
entuser "github.com/cloudreve/Cloudreve/v4/ent/user"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory/types"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/boolset"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/conf"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIsValidShareChecksOwnerAccess(t *testing.T) {
|
||||||
|
permissions := &boolset.BooleanSet{}
|
||||||
|
boolset.Set(types.GroupPermissionShare, true, permissions)
|
||||||
|
allowedGroup := &ent.Group{Permissions: permissions}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
status entuser.Status
|
||||||
|
group *ent.Group
|
||||||
|
wantErr error
|
||||||
|
}{
|
||||||
|
{name: "active owner with share permission", status: entuser.StatusActive, group: allowedGroup},
|
||||||
|
{name: "active owner without share permission", status: entuser.StatusActive, group: &ent.Group{Permissions: &boolset.BooleanSet{}}, wantErr: ErrSourceFileInvalid},
|
||||||
|
{name: "missing group", status: entuser.StatusActive, wantErr: ErrSourceFileInvalid},
|
||||||
|
{name: "missing permissions", status: entuser.StatusActive, group: &ent.Group{}, wantErr: ErrSourceFileInvalid},
|
||||||
|
{name: "manually banned owner", status: entuser.StatusManualBanned, group: allowedGroup, wantErr: ErrOwnerInactive},
|
||||||
|
{name: "system banned owner", status: entuser.StatusSysBanned, group: allowedGroup, wantErr: ErrOwnerInactive},
|
||||||
|
{name: "inactive owner", status: entuser.StatusInactive, group: allowedGroup, wantErr: ErrOwnerInactive},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
owner := &ent.User{ID: 1, Status: tt.status}
|
||||||
|
owner.SetGroup(tt.group)
|
||||||
|
share := &ent.Share{}
|
||||||
|
share.SetUser(owner)
|
||||||
|
share.SetFile(&ent.File{OwnerID: owner.ID, FileChildren: 1})
|
||||||
|
|
||||||
|
require.ErrorIs(t, IsValidShare(share), tt.wantErr)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShareClientRevalidatesOwnerGroup(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()
|
||||||
|
|
||||||
|
permissions := &boolset.BooleanSet{}
|
||||||
|
boolset.Set(types.GroupPermissionShare, true, permissions)
|
||||||
|
group := client.Group.Create().SetName("sharing enabled").SetPermissions(permissions).SaveX(ctx)
|
||||||
|
restrictedGroup := client.Group.Create().SetName("sharing disabled").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
|
||||||
|
owner := client.User.Create().SetEmail("owner@example.com").SetNick("owner").SetGroup(group).SaveX(ctx)
|
||||||
|
root := client.File.Create().SetName(RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(owner).SaveX(ctx)
|
||||||
|
file := client.File.Create().SetName("shared.txt").SetType(int(types.FileTypeFile)).SetOwner(owner).SetParent(root).SaveX(ctx)
|
||||||
|
share := client.Share.Create().SetUser(owner).SetFile(file).SaveX(ctx)
|
||||||
|
shareClient := NewShareClient(client, conf.SQLiteDB, nil)
|
||||||
|
|
||||||
|
// Share-info and listing callers only request the owner and file edges.
|
||||||
|
ctx = context.WithValue(ctx, LoadShareUser{}, true)
|
||||||
|
ctx = context.WithValue(ctx, LoadShareFile{}, true)
|
||||||
|
checkShare := func(wantErr error) {
|
||||||
|
t.Helper()
|
||||||
|
current, err := shareClient.GetByID(ctx, share.ID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.ErrorIs(t, IsValidShare(current), wantErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
checkShare(nil)
|
||||||
|
client.Group.UpdateOne(group).SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
|
||||||
|
checkShare(ErrSourceFileInvalid)
|
||||||
|
client.Group.UpdateOne(group).SetPermissions(permissions).SaveX(ctx)
|
||||||
|
checkShare(nil)
|
||||||
|
client.User.UpdateOne(owner).SetGroup(restrictedGroup).SaveX(ctx)
|
||||||
|
checkShare(ErrSourceFileInvalid)
|
||||||
|
client.User.UpdateOne(owner).SetGroup(group).SaveX(ctx)
|
||||||
|
checkShare(nil)
|
||||||
|
}
|
||||||
@ -0,0 +1,89 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rsa"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"math/big"
|
||||||
|
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
)
|
||||||
|
|
||||||
|
type OIDCIDTokenClaims struct {
|
||||||
|
jwt.RegisteredClaims
|
||||||
|
Nonce string `json:"nonce,omitempty"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
PreferredUsername string `json:"preferred_username,omitempty"`
|
||||||
|
Picture string `json:"picture,omitempty"`
|
||||||
|
UpdatedAt int64 `json:"updated_at,omitempty"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
EmailVerified bool `json:"email_verified,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type JWKSet struct {
|
||||||
|
Keys []JWK `json:"keys"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type JWK struct {
|
||||||
|
Kty string `json:"kty"`
|
||||||
|
Use string `json:"use"`
|
||||||
|
Alg string `json:"alg"`
|
||||||
|
Kid string `json:"kid"`
|
||||||
|
N string `json:"n"`
|
||||||
|
E string `json:"e"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func SignOIDCIDToken(privateKeyRaw string, claims *OIDCIDTokenClaims) (string, error) {
|
||||||
|
key, err := parseRSAPrivateKey(privateKeyRaw)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims)
|
||||||
|
token.Header["kid"] = oidcSigningKeyID(&key.PublicKey)
|
||||||
|
return token.SignedString(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func OIDCJWKSet(privateKeyRaw string) (*JWKSet, error) {
|
||||||
|
key, err := parseRSAPrivateKey(privateKeyRaw)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
kid := oidcSigningKeyID(&key.PublicKey)
|
||||||
|
return &JWKSet{Keys: []JWK{buildJWK(&key.PublicKey, kid)}}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseRSAPrivateKey(privateKeyRaw string) (*rsa.PrivateKey, error) {
|
||||||
|
block, _ := pem.Decode([]byte(privateKeyRaw))
|
||||||
|
if block == nil {
|
||||||
|
return nil, fmt.Errorf("invalid OIDC signing key PEM")
|
||||||
|
}
|
||||||
|
|
||||||
|
key, err := x509.ParsePKCS1PrivateKey(block.Bytes)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid OIDC signing key: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func oidcSigningKeyID(key *rsa.PublicKey) string {
|
||||||
|
der, _ := x509.MarshalPKIXPublicKey(key)
|
||||||
|
sum := sha256.Sum256(der)
|
||||||
|
return base64.RawURLEncoding.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildJWK(key *rsa.PublicKey, kid string) JWK {
|
||||||
|
return JWK{
|
||||||
|
Kty: "RSA",
|
||||||
|
Use: "sig",
|
||||||
|
Alg: "RS256",
|
||||||
|
Kid: kid,
|
||||||
|
N: base64.RawURLEncoding.EncodeToString(key.N.Bytes()),
|
||||||
|
E: base64.RawURLEncoding.EncodeToString(big.NewInt(int64(key.E)).Bytes()),
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -1,61 +0,0 @@
|
|||||||
package cache
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSet(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGet(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
|
|
||||||
value, ok := Get("123")
|
|
||||||
asserts.True(ok)
|
|
||||||
asserts.Equal("321", value)
|
|
||||||
|
|
||||||
value, ok = Get("not_exist")
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeletes(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
err := Deletes([]string{"123"}, "")
|
|
||||||
asserts.NoError(err)
|
|
||||||
_, exist := Get("123")
|
|
||||||
asserts.False(exist)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSettings(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("test_1", "1", -1))
|
|
||||||
|
|
||||||
values, missed := GetSettings([]string{"1", "2"}, "test_")
|
|
||||||
asserts.Equal(map[string]string{"1": "1"}, values)
|
|
||||||
asserts.Equal([]string{"2"}, missed)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSetSettings(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
err := SetSettings(map[string]string{"3": "3", "4": "4"}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
value1, _ := Get("test_3")
|
|
||||||
value2, _ := Get("test_4")
|
|
||||||
asserts.Equal("3", value1)
|
|
||||||
asserts.Equal("4", value2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInit(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@ -1,94 +0,0 @@
|
|||||||
package conf
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/cloudreve/Cloudreve/v4/pkg/util"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"io/ioutil"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 测试Init日志路径错误
|
|
||||||
func TestInitPanic(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 日志路径不存在时
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init("not/exist/path")
|
|
||||||
})
|
|
||||||
|
|
||||||
asserts.True(util.Exists("conf.ini"))
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestInitDelimiterNotFound 日志路径存在但 Key 格式错误时
|
|
||||||
func TestInitDelimiterNotFound(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testCase := `[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password233root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
asserts.Panics(func() {
|
|
||||||
Init("testConf.ini")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestInitNoPanic 日志路径存在且合法时
|
|
||||||
func TestInitNoPanic(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testCase := `
|
|
||||||
[System]
|
|
||||||
Listen = 3000
|
|
||||||
HashIDSalt = 1
|
|
||||||
|
|
||||||
[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password = root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init("testConf.ini")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMapSection(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//正常情况
|
|
||||||
testCase := `
|
|
||||||
[System]
|
|
||||||
Listen = 3000
|
|
||||||
HashIDSalt = 1
|
|
||||||
|
|
||||||
[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password:root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
Init("testConf.ini")
|
|
||||||
err = mapSection("Database", DatabaseConfig)
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -0,0 +1,87 @@
|
|||||||
|
package dbfs
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent"
|
||||||
|
entuser "github.com/cloudreve/Cloudreve/v4/ent/user"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory/types"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/setting"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type directLinkFileClient struct {
|
||||||
|
inventory.FileClient
|
||||||
|
root *ent.File
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *directLinkFileClient) GetParentFile(_ context.Context, file *ent.File, _ bool) (*ent.File, error) {
|
||||||
|
if file.FileChildren == c.root.ID {
|
||||||
|
return c.root, nil
|
||||||
|
}
|
||||||
|
return nil, &ent.NotFoundError{}
|
||||||
|
}
|
||||||
|
|
||||||
|
type directLinkSettingProvider struct {
|
||||||
|
setting.Provider
|
||||||
|
}
|
||||||
|
|
||||||
|
func (directLinkSettingProvider) DBFS(context.Context) *setting.DBFS {
|
||||||
|
return &setting.DBFS{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetFileFromDirectLinkChecksOwnerAccess(t *testing.T) {
|
||||||
|
allowedGroup := &ent.Group{Settings: &types.GroupSetting{SourceBatchSize: 1}}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
status entuser.Status
|
||||||
|
group *ent.Group
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{name: "active owner with direct link permission", status: entuser.StatusActive, group: allowedGroup},
|
||||||
|
{name: "active owner without direct link permission", status: entuser.StatusActive, group: &ent.Group{Settings: &types.GroupSetting{}}, wantErr: true},
|
||||||
|
{name: "negative batch size", status: entuser.StatusActive, group: &ent.Group{Settings: &types.GroupSetting{SourceBatchSize: -1}}, wantErr: true},
|
||||||
|
{name: "missing group", status: entuser.StatusActive, wantErr: true},
|
||||||
|
{name: "missing group settings", status: entuser.StatusActive, group: &ent.Group{}, wantErr: true},
|
||||||
|
{name: "manually banned owner", status: entuser.StatusManualBanned, group: allowedGroup, wantErr: true},
|
||||||
|
{name: "system banned owner", status: entuser.StatusSysBanned, group: allowedGroup, wantErr: true},
|
||||||
|
{name: "inactive owner", status: entuser.StatusInactive, group: allowedGroup, wantErr: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
owner := &ent.User{ID: 1, Status: tt.status}
|
||||||
|
owner.SetGroup(tt.group)
|
||||||
|
root := &ent.File{ID: 1, Name: inventory.RootFolderName, OwnerID: owner.ID}
|
||||||
|
file := &ent.File{ID: 2, Name: "shared.txt", OwnerID: owner.ID, FileChildren: root.ID}
|
||||||
|
file.SetOwner(owner)
|
||||||
|
link := &ent.DirectLink{}
|
||||||
|
link.SetFile(file)
|
||||||
|
dbfs := &DBFS{
|
||||||
|
user: owner,
|
||||||
|
fileClient: &directLinkFileClient{root: root},
|
||||||
|
settingClient: directLinkSettingProvider{},
|
||||||
|
}
|
||||||
|
|
||||||
|
got, err := dbfs.GetFileFromDirectLink(context.Background(), link)
|
||||||
|
if got != nil {
|
||||||
|
t.Cleanup(got.(*File).Parent.Recycle)
|
||||||
|
}
|
||||||
|
if tt.wantErr {
|
||||||
|
require.Nil(t, got)
|
||||||
|
var appErr serializer.AppError
|
||||||
|
require.ErrorAs(t, err, &appErr)
|
||||||
|
require.Equal(t, fs.ErrDirectLinkInvalid.Code, appErr.Code)
|
||||||
|
require.Equal(t, fs.ErrDirectLinkInvalid.Msg, appErr.Msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Same(t, file, got.(*File).Model)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -0,0 +1,236 @@
|
|||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// IPMatcher checks whether an IP address matches a given filter.
|
||||||
|
// The filter can be:
|
||||||
|
// - A single IP address (e.g., "192.168.1.1" or "::1")
|
||||||
|
// - A CIDR notation (e.g., "192.168.1.0/24" or "2001:db8::/32")
|
||||||
|
// - An IP range (e.g., "192.168.1.1-192.168.1.255")
|
||||||
|
// - A wildcard pattern (e.g., "192.168.1.*" or "192.168.*.*")
|
||||||
|
type IPMatcher struct {
|
||||||
|
cidr *net.IPNet
|
||||||
|
ipStart net.IP
|
||||||
|
ipEnd net.IP
|
||||||
|
exact net.IP
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewIPMatcher creates an IPMatcher from a filter string.
|
||||||
|
// Supported formats:
|
||||||
|
// - CIDR: "192.168.1.0/24", "2001:db8::/32"
|
||||||
|
// - Range: "192.168.1.1-192.168.1.255", "::1-::10"
|
||||||
|
// - Wildcard: "192.168.1.*", "192.168.*.*", "10.*.*.*"
|
||||||
|
// - Single IP: "192.168.1.1", "::1"
|
||||||
|
func NewIPMatcher(filter string) (*IPMatcher, error) {
|
||||||
|
filter = strings.TrimSpace(filter)
|
||||||
|
if filter == "" {
|
||||||
|
return nil, fmt.Errorf("empty IP filter")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try CIDR notation first
|
||||||
|
if strings.Contains(filter, "/") {
|
||||||
|
_, cidrNet, err := net.ParseCIDR(filter)
|
||||||
|
if err == nil {
|
||||||
|
return &IPMatcher{cidr: cidrNet}, nil
|
||||||
|
}
|
||||||
|
// If it has "/" but isn't valid CIDR, continue to try other formats
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try IP range (e.g., "192.168.1.1-192.168.1.255")
|
||||||
|
if strings.Count(filter, "-") == 1 {
|
||||||
|
parts := strings.SplitN(filter, "-", 2)
|
||||||
|
start := net.ParseIP(strings.TrimSpace(parts[0]))
|
||||||
|
end := net.ParseIP(strings.TrimSpace(parts[1]))
|
||||||
|
if start != nil && end != nil {
|
||||||
|
// Ensure same IP version
|
||||||
|
if (start.To4() != nil) == (end.To4() != nil) {
|
||||||
|
return &IPMatcher{ipStart: start, ipEnd: end}, nil
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("IP range must contain same IP version: %s", filter)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try wildcard pattern (e.g., "192.168.1.*" or "192.168.*.*")
|
||||||
|
if strings.Contains(filter, "*") {
|
||||||
|
cidr, err := wildcardToCIDR(filter)
|
||||||
|
if err == nil {
|
||||||
|
return &IPMatcher{cidr: cidr}, nil
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try single exact IP
|
||||||
|
ip := net.ParseIP(filter)
|
||||||
|
if ip != nil {
|
||||||
|
return &IPMatcher{exact: ip}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("invalid IP filter format: %s", filter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match checks if the given IP address matches the filter.
|
||||||
|
// The ip parameter should be a string representation of an IP address.
|
||||||
|
// Returns an error if the IP string is invalid.
|
||||||
|
func (m *IPMatcher) Match(ip string) (bool, error) {
|
||||||
|
parsedIP := net.ParseIP(strings.TrimSpace(ip))
|
||||||
|
if parsedIP == nil {
|
||||||
|
return false, fmt.Errorf("invalid IP address: %s", ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
return m.MatchIP(parsedIP), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchIP checks if the given parsed IP address matches the filter.
|
||||||
|
func (m *IPMatcher) MatchIP(ip net.IP) bool {
|
||||||
|
if ip == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case m.cidr != nil:
|
||||||
|
return m.cidr.Contains(ip)
|
||||||
|
case m.ipStart != nil && m.ipEnd != nil:
|
||||||
|
return ipInRange(ip, m.ipStart, m.ipEnd)
|
||||||
|
case m.exact != nil:
|
||||||
|
return m.exact.Equal(ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipInRange checks if ip is within the inclusive range [start, end].
|
||||||
|
func ipInRange(ip, start, end net.IP) bool {
|
||||||
|
// Normalize to 16-byte representation for comparison
|
||||||
|
ip16 := ip.To16()
|
||||||
|
start16 := start.To16()
|
||||||
|
end16 := end.To16()
|
||||||
|
|
||||||
|
return bytesCompare(start16, ip16) <= 0 && bytesCompare(ip16, end16) <= 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// bytesCompare compares two byte slices lexicographically.
|
||||||
|
// Returns -1 if a < b, 0 if a == b, 1 if a > b.
|
||||||
|
func bytesCompare(a, b []byte) int {
|
||||||
|
for i := 0; i < len(a) && i < len(b); i++ {
|
||||||
|
if a[i] < b[i] {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
if a[i] > b[i] {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(a) < len(b) {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
if len(a) > len(b) {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// wildcardToCIDR converts a wildcard IP pattern to a CIDR net.
|
||||||
|
// Only supports IPv4 wildcards (e.g., "192.168.1.*" -> "192.168.1.0/24").
|
||||||
|
func wildcardToCIDR(pattern string) (*net.IPNet, error) {
|
||||||
|
pattern = strings.TrimSpace(pattern)
|
||||||
|
parts := strings.Split(pattern, ".")
|
||||||
|
|
||||||
|
if len(parts) != 4 {
|
||||||
|
return nil, fmt.Errorf("wildcard pattern currently only supports IPv4: %s", pattern)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Count how many fixed octets
|
||||||
|
fixedOctets := 0
|
||||||
|
for i, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if part == "*" {
|
||||||
|
// Fill remaining octets with 0 for the network address
|
||||||
|
for j := i; j < 4; j++ {
|
||||||
|
parts[j] = "0"
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
fixedOctets++
|
||||||
|
}
|
||||||
|
|
||||||
|
// All wildcards would be 0.0.0.0/0
|
||||||
|
if fixedOctets == 0 {
|
||||||
|
return &net.IPNet{
|
||||||
|
IP: net.IPv4(0, 0, 0, 0),
|
||||||
|
Mask: net.CIDRMask(0, 32),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Construct CIDR network
|
||||||
|
cidrStr := fmt.Sprintf("%s/%d", strings.Join(parts, "."), fixedOctets*8)
|
||||||
|
_, cidrNet, err := net.ParseCIDR(cidrStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid wildcard pattern: %s (%w)", pattern, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return cidrNet, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsCIDRNotation checks if a filter string is a valid CIDR notation.
|
||||||
|
func IsCIDRNotation(filter string) bool {
|
||||||
|
_, _, err := net.ParseCIDR(strings.TrimSpace(filter))
|
||||||
|
return err == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPFilterToCIDRs converts an IP filter string to a list of CIDR networks.
|
||||||
|
// This is useful when you need to convert user-friendly IP filters
|
||||||
|
// to CIDR notation for display or storage.
|
||||||
|
// Supports CIDR, wildcard, single IP, and IP range (expanded).
|
||||||
|
func IPFilterToCIDRs(filter string) ([]*net.IPNet, error) {
|
||||||
|
matcher, err := NewIPMatcher(filter)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case matcher.cidr != nil:
|
||||||
|
return []*net.IPNet{matcher.cidr}, nil
|
||||||
|
case matcher.exact != nil:
|
||||||
|
// Single IP -> /32 or /128
|
||||||
|
if matcher.exact.To4() != nil {
|
||||||
|
return []*net.IPNet{{
|
||||||
|
IP: matcher.exact,
|
||||||
|
Mask: net.CIDRMask(32, 32),
|
||||||
|
}}, nil
|
||||||
|
}
|
||||||
|
return []*net.IPNet{{
|
||||||
|
IP: matcher.exact,
|
||||||
|
Mask: net.CIDRMask(128, 128),
|
||||||
|
}}, nil
|
||||||
|
case matcher.ipStart != nil && matcher.ipEnd != nil:
|
||||||
|
// IP range - return as two /32 or /128 CIDRs
|
||||||
|
// Note: This is a simplified representation; the caller should
|
||||||
|
// use Match/MatchIP for exact range matching
|
||||||
|
cidrs := make([]*net.IPNet, 0)
|
||||||
|
if matcher.ipStart.To4() != nil {
|
||||||
|
cidrs = append(cidrs, &net.IPNet{
|
||||||
|
IP: matcher.ipStart,
|
||||||
|
Mask: net.CIDRMask(32, 32),
|
||||||
|
})
|
||||||
|
cidrs = append(cidrs, &net.IPNet{
|
||||||
|
IP: matcher.ipEnd,
|
||||||
|
Mask: net.CIDRMask(32, 32),
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
cidrs = append(cidrs, &net.IPNet{
|
||||||
|
IP: matcher.ipStart,
|
||||||
|
Mask: net.CIDRMask(128, 128),
|
||||||
|
})
|
||||||
|
cidrs = append(cidrs, &net.IPNet{
|
||||||
|
IP: matcher.ipEnd,
|
||||||
|
Mask: net.CIDRMask(128, 128),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return cidrs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("cannot convert filter to CIDR: %s", filter)
|
||||||
|
}
|
||||||
@ -0,0 +1,432 @@
|
|||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewIPMatcher_CIDR(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"IPv4 CIDR /24", "192.168.1.0/24", false},
|
||||||
|
{"IPv4 CIDR /16", "10.0.0.0/16", false},
|
||||||
|
{"IPv4 CIDR /32", "192.168.1.1/32", false},
|
||||||
|
{"IPv4 CIDR /0", "0.0.0.0/0", false},
|
||||||
|
{"IPv6 CIDR /64", "2001:db8::/64", false},
|
||||||
|
{"IPv6 CIDR /128", "::1/128", false},
|
||||||
|
{"Invalid CIDR - bad mask", "192.168.1.0/33", true},
|
||||||
|
{"Invalid CIDR - bad IP", "999.999.999.999/24", true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher(tt.filter)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected error, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) unexpected error: %v", tt.filter, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matcher.cidr == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected CIDR matcher, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewIPMatcher_Range(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"IPv4 range", "192.168.1.1-192.168.1.255", false},
|
||||||
|
{"IPv4 range with spaces", "192.168.1.1 - 192.168.1.255", false},
|
||||||
|
{"IPv6 range", "::1-::10", false},
|
||||||
|
{"Cross IP version", "192.168.1.1-::1", true},
|
||||||
|
{"Invalid start", "abc-192.168.1.255", true},
|
||||||
|
{"Invalid end", "192.168.1.1-abc", true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher(tt.filter)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected error, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) unexpected error: %v", tt.filter, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matcher.ipStart == nil || matcher.ipEnd == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected range matcher, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewIPMatcher_Wildcard(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"One wildcard", "192.168.1.*", false},
|
||||||
|
{"Two wildcards", "192.168.*.*", false},
|
||||||
|
{"Three wildcards", "192.*.*.*", false},
|
||||||
|
{"All wildcards", "*.*.*.*", false},
|
||||||
|
{"IPv6 not supported", "2001:db8::*", true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher(tt.filter)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected error, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) unexpected error: %v", tt.filter, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matcher.cidr == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected CIDR matcher (from wildcard), got nil", tt.filter)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewIPMatcher_SingleIP(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"IPv4 single", "192.168.1.1", false},
|
||||||
|
{"IPv6 single", "::1", false},
|
||||||
|
{"IPv6 full", "2001:db8::1", false},
|
||||||
|
{"Invalid", "not-an-ip", true},
|
||||||
|
{"Empty", "", true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher(tt.filter)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected error, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) unexpected error: %v", tt.filter, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matcher.exact == nil {
|
||||||
|
t.Errorf("NewIPMatcher(%q) expected exact IP matcher, got nil", tt.filter)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_Match_CIDR(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("192.168.1.0/24")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.1", true},
|
||||||
|
{"192.168.1.255", true},
|
||||||
|
{"192.168.1.0", true},
|
||||||
|
{"192.168.2.1", false},
|
||||||
|
{"10.0.0.1", false},
|
||||||
|
{"invalid", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, err := matcher.Match(tt.ip)
|
||||||
|
if tt.ip == "invalid" {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("Match(%q) expected error", tt.ip)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("Match(%q) unexpected error: %v", tt.ip, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_Match_IPv6CIDR(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("2001:db8::/32")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"2001:db8::1", true},
|
||||||
|
{"2001:db8:1234::1", true},
|
||||||
|
{"2001:db9::1", false},
|
||||||
|
{"::1", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, err := matcher.Match(tt.ip)
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("Match(%q) unexpected error: %v", tt.ip, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_Match_Range(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("192.168.1.10-192.168.1.20")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.10", true},
|
||||||
|
{"192.168.1.15", true},
|
||||||
|
{"192.168.1.20", true},
|
||||||
|
{"192.168.1.9", false},
|
||||||
|
{"192.168.1.21", false},
|
||||||
|
{"192.168.2.1", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, err := matcher.Match(tt.ip)
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("Match(%q) unexpected error: %v", tt.ip, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_Match_Wildcard(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("192.168.*.*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.1", true},
|
||||||
|
{"192.168.255.255", true},
|
||||||
|
{"192.168.0.0", true},
|
||||||
|
{"192.169.1.1", false},
|
||||||
|
{"10.0.0.1", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, err := matcher.Match(tt.ip)
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("Match(%q) unexpected error: %v", tt.ip, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_Match_SingleIP(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("192.168.1.1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.1", true},
|
||||||
|
{"192.168.1.2", false},
|
||||||
|
{"192.168.1.0", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, _ := matcher.Match(tt.ip)
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_MatchIP_WithNetIP(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("10.0.0.0/8")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test MatchIP with pre-parsed net.IP
|
||||||
|
ip := net.ParseIP("10.255.255.255")
|
||||||
|
if !matcher.MatchIP(ip) {
|
||||||
|
t.Errorf("MatchIP(%v) = false, want true", ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
ip = net.ParseIP("11.0.0.1")
|
||||||
|
if matcher.MatchIP(ip) {
|
||||||
|
t.Errorf("MatchIP(%v) = true, want false", ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test nil IP
|
||||||
|
if matcher.MatchIP(nil) {
|
||||||
|
t.Errorf("MatchIP(nil) = true, want false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWildcardToCIDR(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
pattern string
|
||||||
|
cidr string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.*", "192.168.1.0/24", false},
|
||||||
|
{"192.168.*.*", "192.168.0.0/16", false},
|
||||||
|
{"192.*.*.*", "192.0.0.0/8", false},
|
||||||
|
{"*.*.*.*", "0.0.0.0/0", false},
|
||||||
|
{"10.0.*.*", "10.0.0.0/16", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.pattern, func(t *testing.T) {
|
||||||
|
cidrNet, err := wildcardToCIDR(tt.pattern)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Errorf("wildcardToCIDR(%q) expected error, got nil", tt.pattern)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("wildcardToCIDR(%q) unexpected error: %v", tt.pattern, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if cidrNet.String() != tt.cidr {
|
||||||
|
t.Errorf("wildcardToCIDR(%q) = %q, want %q", tt.pattern, cidrNet.String(), tt.cidr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsCIDRNotation(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
filter string
|
||||||
|
isCIDR bool
|
||||||
|
}{
|
||||||
|
{"192.168.1.0/24", true},
|
||||||
|
{"2001:db8::/32", true},
|
||||||
|
{"0.0.0.0/0", true},
|
||||||
|
{"192.168.1.1", false},
|
||||||
|
{"not-a-cidr", false},
|
||||||
|
{"192.168.1.*", false},
|
||||||
|
{"192.168.1.1-192.168.1.255", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.filter, func(t *testing.T) {
|
||||||
|
if got := IsCIDRNotation(tt.filter); got != tt.isCIDR {
|
||||||
|
t.Errorf("IsCIDRNotation(%q) = %v, want %v", tt.filter, got, tt.isCIDR)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_BackwardCompatible(t *testing.T) {
|
||||||
|
// Ensures existing exact IP filtering still works
|
||||||
|
matcher, err := NewIPMatcher("123.45.67.89")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher for exact IP: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
matched, err := matcher.Match("123.45.67.89")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Match failed: %v", err)
|
||||||
|
}
|
||||||
|
if !matched {
|
||||||
|
t.Error("Exact IP should match")
|
||||||
|
}
|
||||||
|
|
||||||
|
matched, err = matcher.Match("123.45.67.90")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Match failed: %v", err)
|
||||||
|
}
|
||||||
|
if matched {
|
||||||
|
t.Error("Different exact IP should not match")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIPMatcher_IPv6Range(t *testing.T) {
|
||||||
|
matcher, err := NewIPMatcher("::1-::5")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create matcher: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
match bool
|
||||||
|
}{
|
||||||
|
{"::1", true},
|
||||||
|
{"::3", true},
|
||||||
|
{"::5", true},
|
||||||
|
{"::6", false},
|
||||||
|
{"::0", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.ip, func(t *testing.T) {
|
||||||
|
matched, err := matcher.Match(tt.ip)
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("Match(%q) unexpected error: %v", tt.ip, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if matched != tt.match {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", tt.ip, matched, tt.match)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -0,0 +1,72 @@
|
|||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/url"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/application/constants"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/application/dependency"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory/types"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/auth"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/cluster/routes"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
type DiscoveryService struct{}
|
||||||
|
|
||||||
|
type JWKService struct{}
|
||||||
|
|
||||||
|
func (s *DiscoveryService) Get(c *gin.Context) *DiscoveryResponse {
|
||||||
|
issuer := oidcIssuer(c)
|
||||||
|
return &DiscoveryResponse{
|
||||||
|
Issuer: issuer.String(),
|
||||||
|
AuthorizationEndpoint: routes.MasterOIDCEndpointUrl(issuer, "/session/authorize"),
|
||||||
|
TokenEndpoint: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/token"),
|
||||||
|
UserInfoEndpoint: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/userinfo"),
|
||||||
|
JWKSURI: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/jwks"),
|
||||||
|
ResponseTypesSupported: []string{
|
||||||
|
"code",
|
||||||
|
},
|
||||||
|
GrantTypesSupported: []string{
|
||||||
|
"authorization_code",
|
||||||
|
},
|
||||||
|
SubjectTypesSupported: []string{
|
||||||
|
"public",
|
||||||
|
},
|
||||||
|
IDTokenSigningAlgValuesSupported: []string{
|
||||||
|
"RS256",
|
||||||
|
},
|
||||||
|
TokenEndpointAuthMethods: []string{
|
||||||
|
"client_secret_post",
|
||||||
|
},
|
||||||
|
CodeChallengeMethodsSupported: []string{
|
||||||
|
"S256",
|
||||||
|
},
|
||||||
|
ScopesSupported: []string{
|
||||||
|
types.ScopeOpenID,
|
||||||
|
types.ScopeProfile,
|
||||||
|
types.ScopeEmail,
|
||||||
|
},
|
||||||
|
ClaimsSupported: []string{
|
||||||
|
"sub",
|
||||||
|
"name",
|
||||||
|
"preferred_username",
|
||||||
|
"picture",
|
||||||
|
"updated_at",
|
||||||
|
"email",
|
||||||
|
"email_verified",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *JWKService) Get(c *gin.Context) (*auth.JWKSet, error) {
|
||||||
|
dep := dependency.FromContext(c)
|
||||||
|
return auth.OIDCJWKSet(dep.SettingProvider().OIDCSigningPrivateKey(c))
|
||||||
|
}
|
||||||
|
|
||||||
|
func oidcIssuer(c *gin.Context) *url.URL {
|
||||||
|
dep := dependency.FromContext(c)
|
||||||
|
issuer := *dep.SettingProvider().SiteURL(c)
|
||||||
|
issuer.RawQuery = ""
|
||||||
|
issuer.Fragment = ""
|
||||||
|
return &issuer
|
||||||
|
}
|
||||||
Loading…
Reference in new issue