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
20 changes: 3 additions & 17 deletions api/stream/stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/gotify/server/v3/auth"
"github.com/gotify/server/v3/config"
"github.com/gotify/server/v3/model"
)

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

for _, allowedOrigin := range allowedOrigins {
if allowedOrigin.MatchString(strings.ToLower(u.Hostname())) {
return true
}
}

return false
return config.MatchesFully(allowedOrigins, strings.ToLower(u.Hostname()))
}

func newUpgrader(allowedWebSocketOrigins []string) *websocket.Upgrader {
compiledAllowedOrigins := compileAllowedWebSocketOrigins(allowedWebSocketOrigins)
compiledAllowedOrigins := config.CompileAllowedOrigins(allowedWebSocketOrigins)
return &websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
Expand All @@ -220,12 +215,3 @@ func newUpgrader(allowedWebSocketOrigins []string) *websocket.Upgrader {
},
}
}

func compileAllowedWebSocketOrigins(allowedOrigins []string) []*regexp.Regexp {
var compiledAllowedOrigins []*regexp.Regexp
for _, origin := range allowedOrigins {
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile(origin))
}

return compiledAllowedOrigins
}
20 changes: 11 additions & 9 deletions api/stream/stream_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import (
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/gotify/server/v3/auth"
"github.com/gotify/server/v3/config"
"github.com/gotify/server/v3/mode"
"github.com/gotify/server/v3/model"
"github.com/stretchr/testify/assert"
Expand Down Expand Up @@ -481,16 +482,22 @@ func Test_isAllowedOrigin_withoutAllowedOrigins_failsWhenNotSameOrigin(t *testin

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

req := httptest.NewRequest("GET", "http://example.me/stream", nil)
req.Header.Set("Origin", "http://gorify.example.com")
req.Header.Set("Origin", "http://gotify.net")
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))

req.Header.Set("Origin", "http://push.gotify.net")
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))

req.Header.Set("Origin", "http://go.example.com")
req.Header.Set("Origin", "http://other.gotify.net")
assert.True(t, isAllowedOrigin(req, compiledAllowedOrigins))

req.Header.Set("Origin", "http://hello.example.com")
req.Header.Set("Origin", "http://gotify.net.evil.net")
assert.False(t, isAllowedOrigin(req, compiledAllowedOrigins))

req.Header.Set("Origin", "http://evil-gotify.net")
assert.False(t, isAllowedOrigin(req, compiledAllowedOrigins))
}

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

func Test_compileAllowedWebSocketOrigins(t *testing.T) {
assert.Equal(t, 0, len(compileAllowedWebSocketOrigins([]string{})))
assert.Equal(t, 3, len(compileAllowedWebSocketOrigins([]string{"^.*$", "", "abc"})))
}

func clients(api *API, user uint) []*client {
api.lock.RLock()
defer api.lock.RUnlock()
Expand Down
19 changes: 2 additions & 17 deletions auth/cors.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package auth

import (
"regexp"
"strings"
"time"

Expand All @@ -15,16 +14,11 @@ func CorsConfig(conf *config.Configuration) cors.Config {
MaxAge: 12 * time.Hour,
AllowBrowserExtensions: true,
}
compiledOrigins := compileAllowedCORSOrigins(conf.Server.Cors.AllowOrigins)
compiledOrigins := config.CompileAllowedOrigins(conf.Server.Cors.AllowOrigins)
corsConf.AllowMethods = conf.Server.Cors.AllowMethods
corsConf.AllowHeaders = conf.Server.Cors.AllowHeaders
corsConf.AllowOriginFunc = func(origin string) bool {
for _, compiledOrigin := range compiledOrigins {
if compiledOrigin.MatchString(strings.ToLower(origin)) {
return true
}
}
return false
return config.MatchesFully(compiledOrigins, strings.ToLower(origin))
}
if allowedOrigin := headerIgnoreCase(conf, "access-control-allow-origin"); allowedOrigin != "" && len(compiledOrigins) == 0 {
corsConf.AllowOrigins = append(corsConf.AllowOrigins, allowedOrigin)
Expand All @@ -41,12 +35,3 @@ func headerIgnoreCase(conf *config.Configuration, search string) (value string)
}
return ""
}

func compileAllowedCORSOrigins(allowedOrigins []string) []*regexp.Regexp {
var compiledAllowedOrigins []*regexp.Regexp
for _, origin := range allowedOrigins {
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile(origin))
}

return compiledAllowedOrigins
}
10 changes: 6 additions & 4 deletions auth/cors_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ import (
func TestCorsConfig(t *testing.T) {
mode.Set(mode.Prod)
serverConf := config.Configuration{}
serverConf.Server.Cors.AllowOrigins = []string{"http://test.com"}
serverConf.Server.Cors.AllowOrigins = []string{"http://gotify\\.net|http://push\\.gotify\\.net", "http://other\\.gotify\\.net"}
serverConf.Server.Cors.AllowHeaders = []string{"content-type"}
serverConf.Server.Cors.AllowMethods = []string{"GET"}

Expand All @@ -29,9 +29,11 @@ func TestCorsConfig(t *testing.T) {
AllowBrowserExtensions: true,
}, actual)
assert.NotNil(t, allowF)
assert.True(t, allowF("http://test.com"))
assert.False(t, allowF("https://test.com"))
assert.False(t, allowF("https://other.com"))
assert.True(t, allowF("http://gotify.net"))
assert.True(t, allowF("http://push.gotify.net"))
assert.True(t, allowF("http://other.gotify.net"))
assert.False(t, allowF("http://gotify.net.evil.net"))
assert.False(t, allowF("http://evil-gotify.net"))
}

func TestEmptyCorsConfigWithResponseHeaders(t *testing.T) {
Expand Down
23 changes: 23 additions & 0 deletions config/origin.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
package config

import "regexp"

// CompileAllowedOrigins compiles the patterns as fully matching regexes.
func CompileAllowedOrigins(allowedOrigins []string) []*regexp.Regexp {
var compiledAllowedOrigins []*regexp.Regexp
for _, origin := range allowedOrigins {
compiledAllowedOrigins = append(compiledAllowedOrigins, regexp.MustCompile("^(?:"+origin+")$"))
}

return compiledAllowedOrigins
}

// MatchesFully checks if any of the regexes matches the origin.
func MatchesFully(compiledOrigins []*regexp.Regexp, origin string) bool {
for _, compiledOrigin := range compiledOrigins {
if compiledOrigin.MatchString(origin) {
return true
}
}
return false
}
22 changes: 22 additions & 0 deletions config/origin_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
package config

import (
"testing"

"github.com/stretchr/testify/assert"
)

func TestCompileAllowedOrigins(t *testing.T) {
assert.Equal(t, 0, len(CompileAllowedOrigins([]string{})))
assert.Equal(t, 3, len(CompileAllowedOrigins([]string{"^.*$", "", "abc"})))
}

func TestMatchesFully(t *testing.T) {
compiledOrigins := CompileAllowedOrigins([]string{"gotify\\.net|push\\.gotify\\.net", "other\\.gotify\\.net"})

assert.True(t, MatchesFully(compiledOrigins, "gotify.net"))
assert.True(t, MatchesFully(compiledOrigins, "push.gotify.net"))
assert.True(t, MatchesFully(compiledOrigins, "other.gotify.net"))
assert.False(t, MatchesFully(compiledOrigins, "gotify.net.evil.net"))
assert.False(t, MatchesFully(compiledOrigins, "evil-gotify.net"))
}
Loading