mirror of
https://github.com/zhouxiaoka/autoclip.git
synced 2026-10-02 02:34:34 +08:00
fix(security): 其他网页不能再读取本地后端的 API Key 或发起写请求
CORS 由 * + credentials 收紧为 AutoClip 自家来源;带 Origin 的写请求必须同源或在白名单内, 表单/multipart 简单请求也会被拦;桌面模式只接受本机 Host,防 DNS 重绑定。 调试路由默认不挂载,web 入口默认只听 127.0.0.1。Docker 局域网访问用 AUTOCLIP_ALLOWED_ORIGINS。 Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
API v1 package for FastAPI routes.
|
||||
统一管理所有API路由
|
||||
"""
|
||||
import os
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
@@ -48,7 +49,9 @@ api_router.include_router(subtitle_editor_router, prefix="/subtitle-editor", tag
|
||||
api_router.include_router(upload_router, tags=["upload"])
|
||||
api_router.include_router(progress_router, prefix="/progress", tags=["progress"])
|
||||
api_router.include_router(pipeline_control_router, prefix="/pipeline", tags=["pipeline"])
|
||||
api_router.include_router(debug_router, tags=["debug"])
|
||||
# 调试路由能直接往发布通道塞消息,只在显式开启时挂载
|
||||
if os.getenv("AUTOCLIP_ENABLE_DEBUG_ROUTES", "").lower() in ("1", "true", "yes"):
|
||||
api_router.include_router(debug_router, tags=["debug"])
|
||||
api_router.include_router(simple_progress_router, tags=["simple-progress"])
|
||||
# api_router.include_router(environment_router, tags=["environment"]) # 文件不存在,暂时注释
|
||||
api_router.include_router(settings_router, tags=["settings"])
|
||||
|
||||
@@ -14,6 +14,7 @@ from backend.core.database import engine
|
||||
from backend.models.base import Base
|
||||
from backend.core.config import get_logging_config, get_api_key
|
||||
from backend.core.error_middleware import global_exception_handler
|
||||
from backend.core.local_origin_guard import LocalOriginGuard, allowed_origins
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -60,10 +61,13 @@ def create_app(mode: str = "web") -> FastAPI:
|
||||
# 设置应用状态
|
||||
app.state.mode = mode
|
||||
|
||||
# 配置 CORS
|
||||
# 来源守卫在内、CORS 在外:预检由 CORS 应答,其他网页的写请求由守卫拦下。
|
||||
# 以前是 allow_origins=["*"] + allow_credentials,任意网页都能读设置里的 API Key。
|
||||
origins = allowed_origins()
|
||||
app.add_middleware(LocalOriginGuard, origins=origins, enforce_local_host=(mode == "desktop"))
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"], # 生产环境需要配置具体域名
|
||||
allow_origins=origins,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""
|
||||
本地后端的来源守卫。
|
||||
|
||||
后端没有登录态,任何能发 HTTP 请求到它的页面都能读设置里的 API Key、触发导入和发布。
|
||||
浏览器里打开的普通网页也能向 127.0.0.1 发请求,所以要在服务端拦:
|
||||
|
||||
1. CORS 只对 AutoClip 自己的前端放行,其他网页读不到响应。
|
||||
2. 带 Origin 的写请求(POST/PUT/PATCH/DELETE)必须来自允许的来源,或与请求 Host 同源。
|
||||
表单 / multipart 这类“简单请求”不会触发预检,只靠 CORS 拦不住。
|
||||
3. 桌面模式只接受 Host 为 127.0.0.1 / localhost 的请求,挡住 DNS 重绑定。
|
||||
|
||||
CLI、curl、Tauri 的 reqwest 不带 Origin,照常放行。
|
||||
Docker 通过局域网 IP 或自定义域名访问时,用 AUTOCLIP_ALLOWED_ORIGINS 追加来源(逗号分隔)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Iterable
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
# Tauri 2:macOS / Linux 为 tauri://localhost,Windows 为 http(s)://tauri.localhost;3000 是 Vite 开发端口
|
||||
DEFAULT_ALLOWED_ORIGINS = (
|
||||
"tauri://localhost",
|
||||
"http://tauri.localhost",
|
||||
"https://tauri.localhost",
|
||||
"http://localhost:3000",
|
||||
"http://127.0.0.1:3000",
|
||||
)
|
||||
|
||||
LOCAL_HOSTNAMES = {"127.0.0.1", "localhost", "::1", "[::1]"}
|
||||
SAFE_METHODS = {"GET", "HEAD", "OPTIONS"}
|
||||
|
||||
|
||||
def allowed_origins() -> list[str]:
|
||||
extra = [o.strip().rstrip("/") for o in os.getenv("AUTOCLIP_ALLOWED_ORIGINS", "").split(",") if o.strip()]
|
||||
return list(dict.fromkeys([*DEFAULT_ALLOWED_ORIGINS, *extra]))
|
||||
|
||||
|
||||
def _header(scope: Scope, name: bytes) -> str:
|
||||
for key, value in scope.get("headers") or ():
|
||||
if key == name:
|
||||
return value.decode("latin-1")
|
||||
return ""
|
||||
|
||||
|
||||
def _hostname(host: str) -> str:
|
||||
host = host.strip().lower()
|
||||
if host.startswith("["):
|
||||
return host.split("]")[0] + "]"
|
||||
return host.rsplit(":", 1)[0] if ":" in host else host
|
||||
|
||||
|
||||
def is_same_origin(origin: str, host: str) -> bool:
|
||||
"""Origin 的 host:port 与请求 Host 一致(Docker 直接访问 :8000 或反代保留 Host 的情况)。"""
|
||||
if not origin or not host:
|
||||
return False
|
||||
parsed = urlsplit(origin)
|
||||
return bool(parsed.netloc) and parsed.netloc.lower() == host.strip().lower()
|
||||
|
||||
|
||||
class LocalOriginGuard:
|
||||
def __init__(self, app: ASGIApp, origins: Iterable[str], enforce_local_host: bool) -> None:
|
||||
self.app = app
|
||||
self.origins = {o.lower() for o in origins}
|
||||
self.enforce_local_host = enforce_local_host
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] not in ("http", "websocket"):
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
host = _header(scope, b"host")
|
||||
if self.enforce_local_host and host and _hostname(host) not in LOCAL_HOSTNAMES:
|
||||
await self._reject(scope, receive, send, "host_not_allowed")
|
||||
return
|
||||
|
||||
origin = _header(scope, b"origin").rstrip("/").lower()
|
||||
method = scope.get("method", "GET").upper()
|
||||
cross_site_write = scope["type"] == "http" and method not in SAFE_METHODS
|
||||
if origin and (cross_site_write or scope["type"] == "websocket"):
|
||||
forwarded_host = _header(scope, b"x-forwarded-host")
|
||||
trusted = (
|
||||
origin in self.origins
|
||||
or is_same_origin(origin, host)
|
||||
or is_same_origin(origin, forwarded_host)
|
||||
)
|
||||
if not trusted:
|
||||
await self._reject(scope, receive, send, "origin_not_allowed")
|
||||
return
|
||||
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
async def _reject(self, scope: Scope, receive: Receive, send: Send, code: str) -> None:
|
||||
if scope["type"] == "websocket":
|
||||
await send({"type": "websocket.close", "code": 1008})
|
||||
return
|
||||
response = JSONResponse(
|
||||
status_code=403,
|
||||
content={"detail": "请求来源不被允许", "error_code": code},
|
||||
)
|
||||
await response(scope, receive, send)
|
||||
+5
-2
@@ -25,5 +25,8 @@ if __name__ == "__main__":
|
||||
logger.error(f"无效的端口号: {sys.argv[i + 1]}")
|
||||
port = 8000
|
||||
|
||||
logger.info(f"启动服务器,端口: {port}")
|
||||
uvicorn.run(app, host="0.0.0.0", port=port)
|
||||
# 默认只听本机;Docker / 局域网访问显式设 AUTOCLIP_HOST=0.0.0.0(Dockerfile 的 CMD 自带 --host)
|
||||
import os
|
||||
host = os.getenv("AUTOCLIP_HOST", "127.0.0.1")
|
||||
logger.info(f"启动服务器,地址: {host}:{port}")
|
||||
uvicorn.run(app, host=host, port=port)
|
||||
@@ -0,0 +1,77 @@
|
||||
"""本地后端来源守卫:其他网页不能读 Key、不能发写请求,自家前端和 CLI 不受影响。"""
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from backend.core.local_origin_guard import LocalOriginGuard, allowed_origins
|
||||
|
||||
|
||||
def make_client(enforce_local_host: bool, base_url: str = "http://127.0.0.1:8000") -> TestClient:
|
||||
app = FastAPI()
|
||||
origins = allowed_origins()
|
||||
app.add_middleware(LocalOriginGuard, origins=origins, enforce_local_host=enforce_local_host)
|
||||
app.add_middleware(CORSMiddleware, allow_origins=origins, allow_credentials=True,
|
||||
allow_methods=["*"], allow_headers=["*"])
|
||||
|
||||
@app.get("/settings")
|
||||
def read():
|
||||
return {"key": "sk-secret"}
|
||||
|
||||
@app.post("/import")
|
||||
def write():
|
||||
return {"ok": True}
|
||||
|
||||
return TestClient(app, base_url=base_url)
|
||||
|
||||
|
||||
def test_foreign_page_cannot_read_response_via_cors():
|
||||
c = make_client(True)
|
||||
r = c.get("/settings", headers={"Origin": "https://evil.example"})
|
||||
assert "access-control-allow-origin" not in r.headers
|
||||
r = c.get("/settings", headers={"Origin": "tauri://localhost"})
|
||||
assert r.headers["access-control-allow-origin"] == "tauri://localhost"
|
||||
|
||||
|
||||
def test_foreign_page_cannot_post_simple_request():
|
||||
c = make_client(True)
|
||||
r = c.post("/import", data={"url": "x"}, headers={"Origin": "https://evil.example"})
|
||||
assert r.status_code == 403
|
||||
assert r.json()["error_code"] == "origin_not_allowed"
|
||||
r = c.post("/import", headers={"Origin": "null"})
|
||||
assert r.status_code == 403
|
||||
|
||||
|
||||
def test_app_origins_and_non_browser_clients_allowed():
|
||||
c = make_client(True)
|
||||
for origin in ("tauri://localhost", "http://tauri.localhost", "http://localhost:3000"):
|
||||
assert c.post("/import", headers={"Origin": origin}).status_code == 200
|
||||
# CLI / curl / Tauri reqwest 不带 Origin
|
||||
assert c.post("/import").status_code == 200
|
||||
|
||||
|
||||
def test_preflight_from_app_origin_succeeds():
|
||||
c = make_client(True)
|
||||
r = c.options("/import", headers={"Origin": "tauri://localhost", "Access-Control-Request-Method": "POST"})
|
||||
assert r.status_code == 200
|
||||
|
||||
|
||||
def test_desktop_rejects_dns_rebinding_host():
|
||||
c = make_client(True, base_url="http://attacker.example:8000")
|
||||
assert c.get("/settings").status_code == 403
|
||||
assert make_client(True, base_url="http://localhost:51234").get("/settings").status_code == 200
|
||||
|
||||
|
||||
def test_web_mode_same_origin_and_extra_origins(monkeypatch):
|
||||
# Docker 从局域网 IP 直接访问 :8000 → 同源放行
|
||||
c = make_client(False, base_url="http://192.168.1.20:8000")
|
||||
assert c.post("/import", headers={"Origin": "http://192.168.1.20:8000"}).status_code == 200
|
||||
# 前端在 :3000、后端在 :8000 的局域网部署 → 需要显式加入
|
||||
assert c.post("/import", headers={"Origin": "http://192.168.1.20:3000"}).status_code == 403
|
||||
monkeypatch.setenv("AUTOCLIP_ALLOWED_ORIGINS", "http://192.168.1.20:3000")
|
||||
c = make_client(False, base_url="http://192.168.1.20:8000")
|
||||
assert c.post("/import", headers={"Origin": "http://192.168.1.20:3000"}).status_code == 200
|
||||
|
||||
|
||||
def test_debug_routes_not_mounted_by_default():
|
||||
from backend.api.v1 import api_router
|
||||
assert not any(getattr(r, "path", "").startswith("/debug") for r in api_router.routes)
|
||||
@@ -48,6 +48,7 @@ services:
|
||||
- API_GEMINI_API_KEY=${API_GEMINI_API_KEY:-}
|
||||
- API_SILICONFLOW_API_KEY=${API_SILICONFLOW_API_KEY:-}
|
||||
- UPLOAD_POST_API_KEY=${UPLOAD_POST_API_KEY:-}
|
||||
- AUTOCLIP_ALLOWED_ORIGINS=${AUTOCLIP_ALLOWED_ORIGINS:-}
|
||||
- UPLOAD_POST_USER=${UPLOAD_POST_USER:-}
|
||||
- AUTOCLIP_YT_SUBTITLE_LANGS=${AUTOCLIP_YT_SUBTITLE_LANGS:-}
|
||||
depends_on:
|
||||
|
||||
@@ -67,3 +67,7 @@ LOG_FILE=backend.log
|
||||
# 环境配置
|
||||
ENVIRONMENT=development
|
||||
DEBUG=true
|
||||
|
||||
# 允许访问后端的前端来源(逗号分隔)。本机 localhost:3000 与桌面端已默认允许;
|
||||
# 从局域网 IP / 自定义域名打开前端时加上,例如 http://192.168.1.20:3000
|
||||
AUTOCLIP_ALLOWED_ORIGINS=
|
||||
|
||||
Reference in New Issue
Block a user