Skip to content

Commit d02796b

Browse files
authored
Merge pull request #1054 from gotify/match-fully-allowed-regex
fix: match allowed regex fully
2 parents 74b75e9 + 8f2dad9 commit d02796b

6 files changed

Lines changed: 67 additions & 47 deletions

File tree

‎api/stream/stream.go‎

Lines changed: 3 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ import (
1111
"github.com/gin-gonic/gin"
1212
"github.com/gorilla/websocket"
1313
"github.com/gotify/server/v3/auth"
14+
"github.com/gotify/server/v3/config"
1415
"github.com/gotify/server/v3/model"
1516
)
1617

@@ -201,17 +202,11 @@ func isAllowedOrigin(r *http.Request, allowedOrigins []*regexp.Regexp) bool {
201202
return true
202203
}
203204

204-
for _, allowedOrigin := range allowedOrigins {
205-
if allowedOrigin.MatchString(strings.ToLower(u.Hostname())) {
206-
return true
207-
}
208-
}
209-
210-
return false
205+
return config.MatchesFully(allowedOrigins, strings.ToLower(u.Hostname()))
211206
}
212207

213208
func newUpgrader(allowedWebSocketOrigins []string) *websocket.Upgrader {
214-
compiledAllowedOrigins := compileAllowedWebSocketOrigins(allowedWebSocketOrigins)
209+
compiledAllowedOrigins := config.CompileAllowedOrigins(allowedWebSocketOrigins)
215210
return &websocket.Upgrader{
216211
ReadBufferSize: 1024,
217212
WriteBufferSize: 1024,
@@ -220,12 +215,3 @@ func newUpgrader(allowedWebSocketOrigins []string) *websocket.Upgrader {
220215
},
221216
}
222217
}
223-
224-
func compileAllowedWebSocketOrigins(allowedOrigins []string) []*regexp.Regexp {
225-
var compiledAllowedOrigins []*regexp.Regexp
226-
for _, origin := range allowedOrigins {
227-
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile(origin))
228-
}
229-
230-
return compiledAllowedOrigins
231-
}

‎api/stream/stream_test.go‎

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@ import (
1616
"github.com/gin-gonic/gin"
1717
"github.com/gorilla/websocket"
1818
"github.com/gotify/server/v3/auth"
19+
"github.com/gotify/server/v3/config"
1920
"github.com/gotify/server/v3/mode"
2021
"github.com/gotify/server/v3/model"
2122
"github.com/stretchr/testify/assert"
@@ -481,16 +482,22 @@ func Test_isAllowedOrigin_withoutAllowedOrigins_failsWhenNotSameOrigin(t *testin
481482

482483
func Test_isAllowedOriginMatching(t *testing.T) {
483484
mode.Set(mode.Prod)
484-
compiledAllowedOrigins := compileAllowedWebSocketOrigins([]string{"go.{4}\\.example\\.com", "go\\.example\\.com"})
485+
compiledAllowedOrigins := config.CompileAllowedOrigins([]string{"gotify\\.net|push\\.gotify\\.net", "other\\.gotify\\.net"})
485486

486487
req := httptest.NewRequest("GET", "http://example.me/stream", nil)
487-
req.Header.Set("Origin", "http://gorify.example.com")
488+
req.Header.Set("Origin", "http://gotify.net")
489+
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))
490+
491+
req.Header.Set("Origin", "http://push.gotify.net")
488492
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))
489493

490-
req.Header.Set("Origin", "http://go.example.com")
494+
req.Header.Set("Origin", "http://other.gotify.net")
491495
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))
492496

493-
req.Header.Set("Origin", "http://hello.example.com")
497+
req.Header.Set("Origin", "http://gotify.net.evil.net")
498+
assert.False(t, isAllowedOrigin(req, compiledAllowedOrigins))
499+
500+
req.Header.Set("Origin", "http://evil-gotify.net")
494501
assert.False(t, isAllowedOrigin(req, compiledAllowedOrigins))
495502
}
496503

@@ -517,11 +524,6 @@ func Test_invalidOrigin_returnsFalse(t *testing.T) {
517524
assert.False(t, actual)
518525
}
519526

520-
func Test_compileAllowedWebSocketOrigins(t *testing.T) {
521-
assert.Equal(t, 0, len(compileAllowedWebSocketOrigins([]string{})))
522-
assert.Equal(t, 3, len(compileAllowedWebSocketOrigins([]string{"^.*$", "", "abc"})))
523-
}
524-
525527
func clients(api *API, user uint) []*client {
526528
api.lock.RLock()
527529
defer api.lock.RUnlock()

‎auth/cors.go‎

Lines changed: 2 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
package auth
22

33
import (
4-
"regexp"
54
"strings"
65
"time"
76

@@ -15,16 +14,11 @@ func CorsConfig(conf *config.Configuration) cors.Config {
1514
MaxAge: 12 * time.Hour,
1615
AllowBrowserExtensions: true,
1716
}
18-
compiledOrigins := compileAllowedCORSOrigins(conf.Server.Cors.AllowOrigins)
17+
compiledOrigins := config.CompileAllowedOrigins(conf.Server.Cors.AllowOrigins)
1918
corsConf.AllowMethods = conf.Server.Cors.AllowMethods
2019
corsConf.AllowHeaders = conf.Server.Cors.AllowHeaders
2120
corsConf.AllowOriginFunc = func(origin string) bool {
22-
for _, compiledOrigin := range compiledOrigins {
23-
if compiledOrigin.MatchString(strings.ToLower(origin)) {
24-
return true
25-
}
26-
}
27-
return false
21+
return config.MatchesFully(compiledOrigins, strings.ToLower(origin))
2822
}
2923
if allowedOrigin := headerIgnoreCase(conf, "access-control-allow-origin"); allowedOrigin != "" && len(compiledOrigins) == 0 {
3024
corsConf.AllowOrigins = append(corsConf.AllowOrigins, allowedOrigin)
@@ -41,12 +35,3 @@ func headerIgnoreCase(conf *config.Configuration, search string) (value string)
4135
}
4236
return ""
4337
}
44-
45-
func compileAllowedCORSOrigins(allowedOrigins []string) []*regexp.Regexp {
46-
var compiledAllowedOrigins []*regexp.Regexp
47-
for _, origin := range allowedOrigins {
48-
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile(origin))
49-
}
50-
51-
return compiledAllowedOrigins
52-
}

‎auth/cors_test.go‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@ import (
1313
func TestCorsConfig(t *testing.T) {
1414
mode.Set(mode.Prod)
1515
serverConf := config.Configuration{}
16-
serverConf.Server.Cors.AllowOrigins = []string{"http://test.com"}
16+
serverConf.Server.Cors.AllowOrigins = []string{"http://gotify\\.net|http://push\\.gotify\\.net", "http://other\\.gotify\\.net"}
1717
serverConf.Server.Cors.AllowHeaders = []string{"content-type"}
1818
serverConf.Server.Cors.AllowMethods = []string{"GET"}
1919

@@ -29,9 +29,11 @@ func TestCorsConfig(t *testing.T) {
2929
AllowBrowserExtensions: true,
3030
}, actual)
3131
assert.NotNil(t, allowF)
32-
assert.True(t, allowF("http://test.com"))
33-
assert.False(t, allowF("https://test.com"))
34-
assert.False(t, allowF("https://other.com"))
32+
assert.True(t, allowF("http://gotify.net"))
33+
assert.True(t, allowF("http://push.gotify.net"))
34+
assert.True(t, allowF("http://other.gotify.net"))
35+
assert.False(t, allowF("http://gotify.net.evil.net"))
36+
assert.False(t, allowF("http://evil-gotify.net"))
3537
}
3638

3739
func TestEmptyCorsConfigWithResponseHeaders(t *testing.T) {

‎config/origin.go‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
package config
2+
3+
import "regexp"
4+
5+
// CompileAllowedOrigins compiles the patterns as fully matching regexes.
6+
func CompileAllowedOrigins(allowedOrigins []string) []*regexp.Regexp {
7+
var compiledAllowedOrigins []*regexp.Regexp
8+
for _, origin := range allowedOrigins {
9+
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile("^(?:"+origin+")$"))
10+
}
11+
12+
return compiledAllowedOrigins
13+
}
14+
15+
// MatchesFully checks if any of the regexes matches the origin.
16+
func MatchesFully(compiledOrigins []*regexp.Regexp, origin string) bool {
17+
for _, compiledOrigin := range compiledOrigins {
18+
if compiledOrigin.MatchString(origin) {
19+
return true
20+
}
21+
}
22+
return false
23+
}

‎config/origin_test.go‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
package config
2+
3+
import (
4+
"testing"
5+
6+
"github.com/stretchr/testify/assert"
7+
)
8+
9+
func TestCompileAllowedOrigins(t *testing.T) {
10+
assert.Equal(t, 0, len(CompileAllowedOrigins([]string{})))
11+
assert.Equal(t, 3, len(CompileAllowedOrigins([]string{"^.*$", "", "abc"})))
12+
}
13+
14+
func TestMatchesFully(t *testing.T) {
15+
compiledOrigins := CompileAllowedOrigins([]string{"gotify\\.net|push\\.gotify\\.net", "other\\.gotify\\.net"})
16+
17+
assert.True(t, MatchesFully(compiledOrigins, "gotify.net"))
18+
assert.True(t, MatchesFully(compiledOrigins, "push.gotify.net"))
19+
assert.True(t, MatchesFully(compiledOrigins, "other.gotify.net"))
20+
assert.False(t, MatchesFully(compiledOrigins, "gotify.net.evil.net"))
21+
assert.False(t, MatchesFully(compiledOrigins, "evil-gotify.net"))
22+
}

0 commit comments

Comments
 (0)