51 lines
1.7 KiB
Go
51 lines
1.7 KiB
Go
package middleware
|
||
|
||
import (
|
||
"crypto/tls"
|
||
"net/http"
|
||
"net/http/httptest"
|
||
"testing"
|
||
)
|
||
|
||
// TestSameOriginSchemeSameOrigin 校验 host 与 scheme 双比对(评审 round2 P2-11):
|
||
// 仅 host 相同、scheme 不同的 Origin(如对 https 站点的 http://host)必须拒绝。
|
||
func TestSameOriginScheme(t *testing.T) {
|
||
mk := func(origin string, tlsConn bool, behindProxy bool, xfp string) *http.Request {
|
||
r := httptest.NewRequest(http.MethodPost, "http://example.com/api/auth/login", nil)
|
||
if origin != "" {
|
||
r.Header.Set("Origin", origin)
|
||
}
|
||
if tlsConn {
|
||
r.TLS = &tls.ConnectionState{}
|
||
}
|
||
if behindProxy {
|
||
r = WithBehindProxy(r, true)
|
||
}
|
||
if xfp != "" {
|
||
r.Header.Set("X-Forwarded-Proto", xfp)
|
||
}
|
||
return r
|
||
}
|
||
|
||
cases := []struct {
|
||
name string
|
||
r *http.Request
|
||
want bool
|
||
}{
|
||
{"http 直连 + http Origin", mk("http://example.com", false, false, ""), true},
|
||
{"http 直连 + https Origin(scheme 不匹配)", mk("https://example.com", false, false, ""), false},
|
||
{"https 直连 + https Origin", mk("https://example.com", true, false, ""), true},
|
||
{"https 直连 + http Origin(scheme 不匹配)", mk("http://example.com", true, false, ""), false},
|
||
{"反代 XFP=https + https Origin", mk("https://example.com", false, true, "https"), true},
|
||
{"反代 XFP=https + http Origin(scheme 不匹配)", mk("http://example.com", false, true, "https"), false},
|
||
{"反代未带 XFP 视为 http + http Origin", mk("http://example.com", false, true, ""), true},
|
||
{"host 不同", mk("http://evil.com", false, false, ""), false},
|
||
{"缺 Origin/Referer", mk("", false, false, ""), false},
|
||
}
|
||
for _, tc := range cases {
|
||
if got := sameOrigin(tc.r); got != tc.want {
|
||
t.Errorf("%s: 期望 %v,实际 %v", tc.name, tc.want, got)
|
||
}
|
||
}
|
||
}
|