feat: per-link IP whitelist, clipboard fallback, and link editing
Build and Push to GHCR / build-and-push (push) Has been cancelled
Build and Push to GHCR / build-and-push (push) Has been cancelled
- Add `allowed_ips` to Link model with IPv4/IPv6 and CIDR support - Validate client IP in proxy auth middleware against link whitelist - Extract client IP from X-Forwarded-For / X-Real-Ip headers - Fix copy button for non-HTTPS contexts via execCommand fallback - Allow editing existing links (name, type, auth mode, rate limit, IPs) - Add dedicated IP whitelist modal for quick editing via [HAPI](https://hapi.run) Co-Authored-By: HAPI <[email protected]>
This commit is contained in:
+34
-17
@@ -14,6 +14,17 @@ import (
|
||||
"mirror-proxy/internal/config"
|
||||
)
|
||||
|
||||
func normalizeIPs(ips []string) []string {
|
||||
var out []string
|
||||
for _, ip := range ips {
|
||||
ip = strings.TrimSpace(ip)
|
||||
if ip != "" {
|
||||
out = append(out, ip)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
type Handler struct {
|
||||
cfg *config.Config
|
||||
onRestart func()
|
||||
@@ -192,10 +203,11 @@ func (h *Handler) GetLinks(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
func (h *Handler) CreateLink(w http.ResponseWriter, r *http.Request) {
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
RateLimit int `json:"rate_limit"`
|
||||
AuthMode string `json:"auth_mode"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
RateLimit int `json:"rate_limit"`
|
||||
AuthMode string `json:"auth_mode"`
|
||||
AllowedIPs []string `json:"allowed_ips"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
@@ -214,14 +226,15 @@ func (h *Handler) CreateLink(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
link := &config.Link{
|
||||
ID: generateID(),
|
||||
Name: req.Name,
|
||||
Token: auth.GenerateToken(),
|
||||
Type: req.Type,
|
||||
AuthMode: req.AuthMode,
|
||||
Enabled: true,
|
||||
RateLimit: req.RateLimit,
|
||||
CreatedAt: time.Now().Unix(),
|
||||
ID: generateID(),
|
||||
Name: req.Name,
|
||||
Token: auth.GenerateToken(),
|
||||
Type: req.Type,
|
||||
AuthMode: req.AuthMode,
|
||||
Enabled: true,
|
||||
RateLimit: req.RateLimit,
|
||||
AllowedIPs: normalizeIPs(req.AllowedIPs),
|
||||
CreatedAt: time.Now().Unix(),
|
||||
}
|
||||
|
||||
h.cfg.SetLink(link)
|
||||
@@ -244,11 +257,12 @@ func (h *Handler) UpdateLink(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled *bool `json:"enabled,omitempty"`
|
||||
RateLimit int `json:"rate_limit"`
|
||||
AuthMode string `json:"auth_mode"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled *bool `json:"enabled,omitempty"`
|
||||
RateLimit int `json:"rate_limit"`
|
||||
AuthMode string `json:"auth_mode"`
|
||||
AllowedIPs []string `json:"allowed_ips,omitempty"`
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
@@ -271,6 +285,9 @@ func (h *Handler) UpdateLink(w http.ResponseWriter, r *http.Request) {
|
||||
if req.AuthMode != "" {
|
||||
link.AuthMode = req.AuthMode
|
||||
}
|
||||
if req.AllowedIPs != nil {
|
||||
link.AllowedIPs = normalizeIPs(req.AllowedIPs)
|
||||
}
|
||||
|
||||
h.cfg.SetLink(link)
|
||||
if err := h.cfg.Save(config.GetConfigFilePath()); err != nil {
|
||||
|
||||
+71
-4
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -82,6 +83,65 @@ func getAuthMode(link *config.Link) string {
|
||||
return link.AuthMode
|
||||
}
|
||||
|
||||
// CheckIPAllowed 检查客户端 IP 是否在链接的白名单中
|
||||
func CheckIPAllowed(clientIP string, allowedIPs []string) bool {
|
||||
if len(allowedIPs) == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
ip := net.ParseIP(clientIP)
|
||||
if ip == nil {
|
||||
// 尝试去掉端口
|
||||
host, _, err := net.SplitHostPort(clientIP)
|
||||
if err == nil {
|
||||
ip = net.ParseIP(host)
|
||||
}
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
for _, allowed := range allowedIPs {
|
||||
allowed = strings.TrimSpace(allowed)
|
||||
if allowed == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// CIDR 格式
|
||||
if strings.Contains(allowed, "/") {
|
||||
_, ipNet, err := net.ParseCIDR(allowed)
|
||||
if err == nil && ipNet.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// 单个 IP
|
||||
allowedIP := net.ParseIP(allowed)
|
||||
if allowedIP != nil && allowedIP.Equal(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// extractClientIP 从请求中提取客户端 IP(去掉端口)
|
||||
func extractClientIP(r *http.Request) string {
|
||||
ip := r.RemoteAddr
|
||||
if xf := r.Header.Get("X-Forwarded-For"); xf != "" {
|
||||
ip = strings.Split(xf, ",")[0]
|
||||
} else if xr := r.Header.Get("X-Real-Ip"); xr != "" {
|
||||
ip = xr
|
||||
}
|
||||
ip = strings.TrimSpace(ip)
|
||||
host, _, err := net.SplitHostPort(ip)
|
||||
if err == nil {
|
||||
ip = host
|
||||
}
|
||||
return ip
|
||||
}
|
||||
|
||||
func ValidateLinkToken(linkID, token string) (*config.Link, bool) {
|
||||
cfg := config.Get()
|
||||
link, ok := cfg.GetLink(linkID)
|
||||
@@ -118,6 +178,7 @@ func ValidateLinkTokenSingle(token string) (*config.Link, bool) {
|
||||
type contextKey string
|
||||
|
||||
const LinkContextKey contextKey = "link"
|
||||
const TokenPrefixContextKey contextKey = "tokenPrefix"
|
||||
|
||||
func ProxyAuthMiddleware(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
@@ -132,11 +193,13 @@ func ProxyAuthMiddleware(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
var link *config.Link
|
||||
var ok bool
|
||||
var newPath string
|
||||
var tokenPrefix string
|
||||
|
||||
// 优先尝试 dual 模式 (/{linkID}/{token}/...)
|
||||
if len(parts) >= 2 {
|
||||
link, ok = ValidateLinkToken(parts[0], parts[1])
|
||||
if ok {
|
||||
tokenPrefix = "/" + parts[0] + "/" + parts[1]
|
||||
if len(parts) == 2 {
|
||||
newPath = "/"
|
||||
} else {
|
||||
@@ -149,6 +212,7 @@ func ProxyAuthMiddleware(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
if !ok {
|
||||
link, ok = ValidateLinkTokenSingle(parts[0])
|
||||
if ok {
|
||||
tokenPrefix = "/" + parts[0]
|
||||
if len(parts) == 1 {
|
||||
newPath = "/"
|
||||
} else {
|
||||
@@ -162,9 +226,11 @@ func ProxyAuthMiddleware(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
return
|
||||
}
|
||||
|
||||
clientIP := r.RemoteAddr
|
||||
if xf := r.Header.Get("X-Forwarded-For"); xf != "" {
|
||||
clientIP = strings.Split(xf, ",")[0]
|
||||
clientIP := extractClientIP(r)
|
||||
|
||||
if !CheckIPAllowed(clientIP, link.AllowedIPs) {
|
||||
http.Error(w, "Forbidden: IP not allowed", http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
if !rl.Allow(clientIP, link.RateLimit) {
|
||||
@@ -174,8 +240,9 @@ func ProxyAuthMiddleware(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
|
||||
r.URL.Path = newPath
|
||||
|
||||
// 存储link到context
|
||||
// 存储link和token前缀到context
|
||||
ctx := context.WithValue(r.Context(), LinkContextKey, link)
|
||||
ctx = context.WithValue(ctx, TokenPrefixContextKey, tokenPrefix)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -39,6 +39,7 @@ type Link struct {
|
||||
AuthMode string `json:"auth_mode"` // single, dual
|
||||
Enabled bool `json:"enabled"`
|
||||
RateLimit int `json:"rate_limit"` // requests per minute
|
||||
AllowedIPs []string `json:"allowed_ips"` // empty = allow all
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
AccessCount int64 `json:"access_count"`
|
||||
LastAccess int64 `json:"last_access"`
|
||||
|
||||
@@ -1,11 +1,16 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"mirror-proxy/internal/auth"
|
||||
)
|
||||
|
||||
// NewGitHubProxy 创建 GitHub 主站反向代理
|
||||
@@ -18,12 +23,50 @@ func NewGitHubProxy() http.Handler {
|
||||
req.URL.Host = target.Host
|
||||
req.Host = target.Host
|
||||
req.Header.Set("Host", target.Host)
|
||||
// 删除 Accept-Encoding,防止响应被压缩,便于修改 HTML
|
||||
req.Header.Del("Accept-Encoding")
|
||||
if req.Header.Get("User-Agent") == "" {
|
||||
req.Header.Set("User-Agent", "MirrorProxy/1.0")
|
||||
}
|
||||
req.Header.Del("X-Forwarded-For")
|
||||
}
|
||||
|
||||
p.ModifyResponse = func(resp *http.Response) error {
|
||||
// 从请求上下文中获取 token 前缀
|
||||
tokenPrefix := ""
|
||||
if resp.Request != nil {
|
||||
if prefix, ok := resp.Request.Context().Value(auth.TokenPrefixContextKey).(string); ok {
|
||||
tokenPrefix = prefix
|
||||
}
|
||||
}
|
||||
|
||||
if tokenPrefix == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 重写 Location header
|
||||
if loc := resp.Header.Get("Location"); loc != "" {
|
||||
resp.Header.Set("Location", rewriteURL(loc, tokenPrefix))
|
||||
}
|
||||
|
||||
// 对 HTML 响应注入 JS 脚本,拦截链接点击
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
if strings.Contains(contentType, "text/html") && resp.Body != nil {
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp.Body.Close()
|
||||
|
||||
body = injectTokenPrefixScript(body, tokenPrefix)
|
||||
resp.Body = io.NopCloser(bytes.NewReader(body))
|
||||
resp.ContentLength = int64(len(body))
|
||||
resp.Header.Set("Content-Length", fmt.Sprintf("%d", len(body)))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
p.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
@@ -40,6 +83,69 @@ func NewGitHubProxy() http.Handler {
|
||||
return p
|
||||
}
|
||||
|
||||
// rewriteURL 重写 URL,在前面加上 token 前缀
|
||||
func rewriteURL(u string, prefix string) string {
|
||||
if u == "" || prefix == "" {
|
||||
return u
|
||||
}
|
||||
// 已经是完整 URL(带 scheme)
|
||||
if strings.HasPrefix(u, "http://") || strings.HasPrefix(u, "https://") {
|
||||
return u
|
||||
}
|
||||
// 已经是带前缀的路径
|
||||
if strings.HasPrefix(u, prefix+"/") || u == prefix {
|
||||
return u
|
||||
}
|
||||
// 以 / 开头的绝对路径,加上前缀
|
||||
if strings.HasPrefix(u, "/") {
|
||||
return prefix + u
|
||||
}
|
||||
// 相对路径,不做处理
|
||||
return u
|
||||
}
|
||||
|
||||
// injectTokenPrefixScript 在 HTML 中注入 JS 脚本,拦截链接点击自动补全 token 前缀
|
||||
func injectTokenPrefixScript(body []byte, prefix string) []byte {
|
||||
script := []byte(`<script>` +
|
||||
`(function(){` +
|
||||
`var p='` + prefix + `';` +
|
||||
`document.addEventListener('click',function(e){` +
|
||||
`var a=e.target.closest('a');` +
|
||||
`if(!a)return;` +
|
||||
`var h=a.getAttribute('href');` +
|
||||
`if(h&&h.startsWith('/')&&!h.startsWith(p+'/')&&h!==p){` +
|
||||
`e.preventDefault();` +
|
||||
`location.href=p+h;` +
|
||||
`}` +
|
||||
`},true);` +
|
||||
`document.addEventListener('submit',function(e){` +
|
||||
`var f=e.target.closest('form');` +
|
||||
`if(!f)return;` +
|
||||
`var h=f.getAttribute('action');` +
|
||||
`if(h&&h.startsWith('/')&&!h.startsWith(p+'/')&&h!==p){` +
|
||||
`f.setAttribute('action',p+h);` +
|
||||
`}` +
|
||||
`},true);` +
|
||||
`})();` +
|
||||
`</script>`)
|
||||
|
||||
// 尝试在 </head> 前插入
|
||||
if idx := bytes.Index(body, []byte("</head>")); idx != -1 {
|
||||
return append(body[:idx], append(script, body[idx:]...)...)
|
||||
}
|
||||
// 或者在 <body> 标签后插入
|
||||
if idx := bytes.Index(body, []byte("<body")); idx != -1 {
|
||||
// 找到 <body> 标签的结束位置
|
||||
endIdx := bytes.Index(body[idx:], []byte(">"))
|
||||
if endIdx != -1 {
|
||||
pos := idx + endIdx + 1
|
||||
return append(body[:pos], append(script, body[pos:]...)...)
|
||||
}
|
||||
}
|
||||
// fallback:在文档开头插入
|
||||
return append(script, body...)
|
||||
}
|
||||
|
||||
// NewGitHubRawProxy 创建 GitHub Raw 反向代理
|
||||
func NewGitHubRawProxy() http.Handler {
|
||||
target, _ := url.Parse("https://raw.githubusercontent.com")
|
||||
@@ -56,6 +162,21 @@ func NewGitHubRawProxy() http.Handler {
|
||||
req.Header.Del("X-Forwarded-For")
|
||||
}
|
||||
|
||||
p.ModifyResponse = func(resp *http.Response) error {
|
||||
tokenPrefix := ""
|
||||
if resp.Request != nil {
|
||||
if prefix, ok := resp.Request.Context().Value(auth.TokenPrefixContextKey).(string); ok {
|
||||
tokenPrefix = prefix
|
||||
}
|
||||
}
|
||||
if tokenPrefix != "" {
|
||||
if loc := resp.Header.Get("Location"); loc != "" {
|
||||
resp.Header.Set("Location", rewriteURL(loc, tokenPrefix))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
p.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
@@ -88,6 +209,21 @@ func NewGitHubAPIProxy() http.Handler {
|
||||
req.Header.Del("X-Forwarded-For")
|
||||
}
|
||||
|
||||
p.ModifyResponse = func(resp *http.Response) error {
|
||||
tokenPrefix := ""
|
||||
if resp.Request != nil {
|
||||
if prefix, ok := resp.Request.Context().Value(auth.TokenPrefixContextKey).(string); ok {
|
||||
tokenPrefix = prefix
|
||||
}
|
||||
}
|
||||
if tokenPrefix != "" {
|
||||
if loc := resp.Header.Get("Location"); loc != "" {
|
||||
resp.Header.Set("Location", rewriteURL(loc, tokenPrefix))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
p.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusBadGateway)
|
||||
|
||||
Reference in New Issue
Block a user