Tweak WebSocket header-processing (#6929)

* fix

* consolidate code
Этот коммит содержится в:
Chris
2017-07-13 14:02:33 -07:00
коммит произвёл Christopher Speller
родитель a1f17c1f84
Коммит 5c3c909c85
3 изменённых файлов: 20 добавлений и 4 удалений

Просмотреть файл

@@ -362,6 +362,15 @@ func TestWebsocketOriginSecurity(t *testing.T) {
t.Fatal("Should have errored because Origin contain AllowCorsFrom") t.Fatal("Should have errored because Origin contain AllowCorsFrom")
} }
// Should fail because non-matching CORS
*utils.Cfg.ServiceSettings.AllowCorsFrom = "http://www.good.com"
_, _, err = websocket.DefaultDialer.Dial(url+model.API_URL_SUFFIX_V3+"/users/websocket", http.Header{
"Origin": []string{"http://www.good.co"},
})
if err == nil {
t.Fatal("Should have errored because Origin does not match host! SECURITY ISSUE!")
}
*utils.Cfg.ServiceSettings.AllowCorsFrom = "" *utils.Cfg.ServiceSettings.AllowCorsFrom = ""
} }

Просмотреть файл

@@ -53,9 +53,8 @@ type CorsWrapper struct {
func (cw *CorsWrapper) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (cw *CorsWrapper) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if len(*utils.Cfg.ServiceSettings.AllowCorsFrom) > 0 { if len(*utils.Cfg.ServiceSettings.AllowCorsFrom) > 0 {
origin := r.Header.Get("Origin") if utils.OriginChecker(r) {
if *utils.Cfg.ServiceSettings.AllowCorsFrom == "*" || strings.Contains(*utils.Cfg.ServiceSettings.AllowCorsFrom, origin) { w.Header().Set("Access-Control-Allow-Origin", r.Header.Get("Origin"))
w.Header().Set("Access-Control-Allow-Origin", origin)
if r.Method == "OPTIONS" { if r.Method == "OPTIONS" {
w.Header().Set( w.Header().Set(

Просмотреть файл

@@ -15,7 +15,15 @@ type OriginCheckerProc func(*http.Request) bool
func OriginChecker(r *http.Request) bool { func OriginChecker(r *http.Request) bool {
origin := r.Header.Get("Origin") origin := r.Header.Get("Origin")
return *Cfg.ServiceSettings.AllowCorsFrom == "*" || strings.Contains(*Cfg.ServiceSettings.AllowCorsFrom, origin) if *Cfg.ServiceSettings.AllowCorsFrom == "*" {
return true
}
for _, allowed := range strings.Split(*Cfg.ServiceSettings.AllowCorsFrom, " ") {
if allowed == origin {
return true
}
}
return false
} }
func GetOriginChecker(r *http.Request) OriginCheckerProc { func GetOriginChecker(r *http.Request) OriginCheckerProc {