diff --git a/internal/middleware/origin_test.go b/internal/middleware/origin_test.go new file mode 100644 index 0000000..74ca5a9 --- /dev/null +++ b/internal/middleware/origin_test.go @@ -0,0 +1,50 @@ +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) + } + } +} diff --git a/internal/middleware/ratelimit.go b/internal/middleware/ratelimit.go index 867ea5f..53a0c36 100644 --- a/internal/middleware/ratelimit.go +++ b/internal/middleware/ratelimit.go @@ -25,7 +25,8 @@ func OriginCheck(next http.Handler) http.Handler { }) } -// sameOrigin 校验 Origin(或 Referer)的 host 与请求 Host 一致。 +// sameOrigin 校验 Origin(或 Referer)与请求同源:host 一致且 scheme 与请求 +// 实际 scheme 一致(TLS 直连为 https;可信反代后取 X-Forwarded-Proto,评审 round2 P2-11)。 func sameOrigin(r *http.Request) bool { raw := r.Header.Get("Origin") if raw == "" { @@ -35,10 +36,16 @@ func sameOrigin(r *http.Request) bool { return false } u, err := url.Parse(raw) - if err != nil || u.Host == "" { + if err != nil || u.Host == "" || u.Scheme == "" { return false } - return strings.EqualFold(u.Host, r.Host) + scheme := "http" + if r.TLS != nil { + scheme = "https" + } else if trueBehindProxy(r) && strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https") { + scheme = "https" + } + return strings.EqualFold(u.Scheme, scheme) && strings.EqualFold(u.Host, r.Host) } func writeErr(w http.ResponseWriter, status int, code, msg string) {