mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-01 15:30:48 +02:00
Co-authored-by: Elias Schneider <login@eliasschneider.com>
This commit is contained in:
co-authored by
Elias Schneider
parent
7c55bdf115
commit
1934efa84c
@@ -123,6 +123,18 @@ func GetCallbackURLFromList(urls []string, inputCallbackURL string) (callbackURL
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// MatchesAnyURLPattern reports whether input matches any pattern in the list,
|
||||
// using the same wildcard rules as callback URLs. An empty list never matches.
|
||||
func MatchesAnyURLPattern(patterns []string, input string) bool {
|
||||
for _, pattern := range patterns {
|
||||
matches, err := matchCallbackURL(pattern, input)
|
||||
if err == nil && matches {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func loopbackURLWithWildcardPort(input string) string {
|
||||
u, _ := url.Parse(input)
|
||||
|
||||
@@ -201,8 +213,14 @@ func normalizeToURLPatternStandard(pattern string) string {
|
||||
var result strings.Builder
|
||||
result.Grow(len(pattern) + 5) // Add 5 for some extra capacity, hoping to avoid many re-allocations
|
||||
|
||||
// First, process the base
|
||||
writeNormalizedBase(&result, patternBase)
|
||||
writeNormalizedPath(&result, patternPath)
|
||||
|
||||
return result.String()
|
||||
}
|
||||
|
||||
// writeNormalizedBase escapes the colons in the scheme and authority that urlpattern would otherwise read as wildcards
|
||||
func writeNormalizedBase(result *strings.Builder, patternBase string) {
|
||||
// 0 = scheme
|
||||
// 1 = hostname (optionally with username/password) - before IPv6 start (no `[` found)
|
||||
// 2 = is matching IPv6 (until `]`)
|
||||
@@ -223,6 +241,12 @@ func normalizeToURLPatternStandard(pattern string) string {
|
||||
case '[':
|
||||
// Start of IPv6 match
|
||||
step = 2
|
||||
case ':':
|
||||
// urlpattern reads ":name" as a single-segment wildcard, but the only wildcards this package supports are * and **
|
||||
// A colon that introduces a port is followed by a digit, so it stays structural and everything else is escaped to a literal
|
||||
if !isPortSeparator(patternBase, i) {
|
||||
result.WriteByte('\\')
|
||||
}
|
||||
}
|
||||
case 2:
|
||||
if patternBase[i] == '/' || patternBase[i] == ']' || patternBase[i] == '[' {
|
||||
@@ -243,8 +267,10 @@ func normalizeToURLPatternStandard(pattern string) string {
|
||||
// Write the byte
|
||||
result.WriteByte(patternBase[i])
|
||||
}
|
||||
}
|
||||
|
||||
// Next, process the path
|
||||
// writeNormalizedPath converts * and ** into the wildcards urlpattern understands, leaving every other character literal
|
||||
func writeNormalizedPath(result *strings.Builder, patternPath string) {
|
||||
for i := 0; i < len(patternPath); i++ {
|
||||
if patternPath[i] == '*' {
|
||||
// Replace globstar with a single asterisk
|
||||
@@ -257,11 +283,19 @@ func normalizeToURLPatternStandard(pattern string) string {
|
||||
result.WriteString(strconv.Itoa(i))
|
||||
}
|
||||
} else {
|
||||
// A literal colon in the path would otherwise be read as a ":name" wildcard
|
||||
if patternPath[i] == ':' {
|
||||
result.WriteByte('\\')
|
||||
}
|
||||
// Add the byte
|
||||
result.WriteByte(patternPath[i])
|
||||
}
|
||||
}
|
||||
return result.String()
|
||||
}
|
||||
|
||||
// isPortSeparator reports whether the colon at index i separates the host from a port
|
||||
func isPortSeparator(s string, i int) bool {
|
||||
return i+1 < len(s) && s[i+1] >= '0' && s[i+1] <= '9'
|
||||
}
|
||||
|
||||
func extractPath(url string) (base string, path string) {
|
||||
|
||||
@@ -699,6 +699,29 @@ func TestGetCallbackURLFromList_LoopbackSpecialHandling(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchesAnyURLPattern(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
patterns []string
|
||||
input string
|
||||
want bool
|
||||
}{
|
||||
{"empty list denies", nil, "https://app.example.com/oauth/client", false},
|
||||
{"empty slice denies", []string{}, "https://app.example.com/oauth/client", false},
|
||||
{"exact match", []string{"https://app.example.com/oauth/client"}, "https://app.example.com/oauth/client", true},
|
||||
{"wildcard path", []string{"https://app.example.com/**"}, "https://app.example.com/oauth/client", true},
|
||||
{"wildcard host segment", []string{"https://*.example.com/oauth/client"}, "https://app.example.com/oauth/client", true},
|
||||
{"star matches all", []string{"*"}, "https://anything.example.com/x", true},
|
||||
{"no match", []string{"https://other.example.com/**"}, "https://app.example.com/oauth/client", false},
|
||||
{"second pattern matches", []string{"https://a.example.com/**", "https://app.example.com/**"}, "https://app.example.com/oauth/client", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, MatchesAnyURLPattern(tt.patterns, tt.input))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoopbackURLWithWildcardPort(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -838,3 +861,42 @@ func TestGetCallbackURLFromList_MultiplePatterns(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The only wildcards this package supports are * and **
|
||||
// urlpattern additionally reads ":name" as a single-segment wildcard, so a literal colon in a
|
||||
// pattern must never widen what it matches
|
||||
func TestMatchCallbackURL_ColonIsNotAWildcard(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
pattern string
|
||||
input string
|
||||
want bool
|
||||
}{
|
||||
{"host label is literal", "https://:host.example.com/cb", "https://evil.example.com/cb", false},
|
||||
{"host label matches itself", "https://:host.example.com/cb", "https://:host.example.com/cb", true},
|
||||
{"path segment is literal", "https://app.example.com/a:b", "https://app.example.com/a:other", false},
|
||||
{"path segment matches itself", "https://app.example.com/a:b", "https://app.example.com/a:b", true},
|
||||
{"userinfo is literal", "https://user:pass@app.example.com/cb", "https://user:other@app.example.com/cb", false},
|
||||
{"userinfo matches itself", "https://user:pass@app.example.com/cb", "https://user:pass@app.example.com/cb", true},
|
||||
|
||||
// Structural colons must keep working
|
||||
{"port is matched exactly", "https://app.example.com:8080/cb", "https://app.example.com:8080/cb", true},
|
||||
{"port mismatch is rejected", "https://app.example.com:8080/cb", "https://app.example.com:9090/cb", false},
|
||||
{"ipv6 host", "https://[::1]/cb", "https://[::1]/cb", true},
|
||||
{"ipv6 host with port", "https://[::1]:8080/cb", "https://[::1]:8080/cb", true},
|
||||
|
||||
// The supported wildcards are unaffected
|
||||
{"single asterisk spans one segment", "https://app.example.com/*/cb", "https://app.example.com/x/cb", true},
|
||||
{"single asterisk does not span two", "https://app.example.com/*/cb", "https://app.example.com/x/y/cb", false},
|
||||
{"globstar spans many segments", "https://app.example.com/**", "https://app.example.com/a/b/c", true},
|
||||
{"asterisk in host", "https://*.example.com/cb", "https://sub.example.com/cb", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := matchCallbackURL(tt.pattern, tt.input)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -28,6 +28,13 @@ var tailscaleIPNets = []*net.IPNet{
|
||||
{IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}, // 100.64.0.0/10
|
||||
}
|
||||
|
||||
// LocalIPv6IPNets returns the extra IPv6 ranges configured via LOCAL_IPV6_RANGES
|
||||
// that are treated as local/private. It is used to extend SSRF protection in
|
||||
// components that classify IPs independently (e.g. the fosite CIMD fetcher).
|
||||
func LocalIPv6IPNets() []*net.IPNet {
|
||||
return localIPv6Ranges
|
||||
}
|
||||
|
||||
func IsLocalIPv6(ip net.IP) bool {
|
||||
if ip.To4() != nil {
|
||||
return false
|
||||
|
||||
@@ -6,12 +6,16 @@ package testing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/italypaleale/francis/components/standalone"
|
||||
"github.com/italypaleale/francis/host/local"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/http3"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -26,8 +30,9 @@ const testActorHostPSK = "pocket-id-test-actor-host-psk-32bytes"
|
||||
func NewActorHostForTest(t *testing.T, register func(t *testing.T, h *local.Host)) *local.Host {
|
||||
t.Helper()
|
||||
|
||||
address := freeLoopbackUDPAddr(t)
|
||||
hostOpts := []local.HostOption{
|
||||
local.WithAddress(freeLoopbackAddr(t)),
|
||||
local.WithAddress(address),
|
||||
local.WithRuntimePSKs([]byte(testActorHostPSK)),
|
||||
local.WithStandaloneMemoryProvider(standalone.StandaloneMemoryOptions{}),
|
||||
local.WithShutdownGracePeriod(time.Second),
|
||||
@@ -62,19 +67,58 @@ func NewActorHostForTest(t *testing.T, register func(t *testing.T, h *local.Host
|
||||
t.Fatal("timed out waiting for the actor host to become ready")
|
||||
}
|
||||
|
||||
// Francis signals host readiness before starting the peer server, so wait for a remote TLS response before a fast test can trigger cleanup
|
||||
waitForActorHostPeerServer(t, address, errCh)
|
||||
|
||||
return h
|
||||
}
|
||||
|
||||
// freeLoopbackAddr reserves a free loopback port and returns its address
|
||||
// waitForActorHostPeerServer waits until the WebTransport listener has passed the startup point that races with shutdown
|
||||
func waitForActorHostPeerServer(t *testing.T, address string, errCh <-chan error) {
|
||||
t.Helper()
|
||||
|
||||
// The probe intentionally omits the Francis client certificate because a remote TLS rejection is enough to prove the peer server is accepting connections
|
||||
//nolint:gosec
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
NextProtos: []string{http3.NextProtoH3},
|
||||
}
|
||||
deadline := time.Now().Add(10 * time.Second)
|
||||
|
||||
for time.Now().Before(deadline) {
|
||||
probeCtx, probeCancel := context.WithTimeout(t.Context(), 200*time.Millisecond)
|
||||
conn, err := quic.DialAddr(probeCtx, address, tlsConfig, &quic.Config{})
|
||||
probeCancel()
|
||||
if conn != nil {
|
||||
_ = conn.CloseWithError(0, "readiness probe complete")
|
||||
return
|
||||
}
|
||||
|
||||
var transportErr *quic.TransportError
|
||||
if errors.As(err, &transportErr) && transportErr.Remote {
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case runErr := <-errCh:
|
||||
t.Fatalf("actor host stopped before its peer server became ready: %v", runErr)
|
||||
case <-time.After(10 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
t.Fatalf("timed out waiting for actor host peer server %s", address)
|
||||
}
|
||||
|
||||
// freeLoopbackUDPAddr reserves a free loopback UDP port and returns its address
|
||||
// The port is released before returning, so the actor host can bind it
|
||||
func freeLoopbackAddr(t *testing.T) string {
|
||||
func freeLoopbackUDPAddr(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
var lc net.ListenConfig
|
||||
lis, err := lc.Listen(t.Context(), "tcp", "127.0.0.1:0")
|
||||
lis, err := lc.ListenPacket(t.Context(), "udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
addr := lis.Addr().String()
|
||||
addr := lis.LocalAddr().String()
|
||||
err = lis.Close()
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user