mirror of
https://github.com/RVC-Boss/GPT-SoVITS.git
synced 2026-07-22 10:03:23 +08:00
Merge 7b8e00b0713af50c6d663a8bcd58e2e5d5b97a8c into 551918539c8e3a496a504a4d9aa2cd7d591097bf
This commit is contained in:
commit
6ce3452417
25
api.py
25
api.py
@ -144,6 +144,7 @@ import argparse
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
_startup_argv = sys.argv[:]
|
||||
|
||||
now_dir = os.getcwd()
|
||||
sys.path.append(now_dir)
|
||||
@ -156,7 +157,7 @@ import torch
|
||||
import torchaudio
|
||||
import librosa
|
||||
import soundfile as sf
|
||||
from fastapi import FastAPI, Request, Query
|
||||
from fastapi import FastAPI, Request, Query, Header, HTTPException, Depends
|
||||
from fastapi.responses import StreamingResponse, JSONResponse
|
||||
import uvicorn
|
||||
from transformers import AutoModelForMaskedLM, AutoTokenizer
|
||||
@ -267,10 +268,10 @@ def init_hifigan():
|
||||
)
|
||||
hifigan_model.eval()
|
||||
hifigan_model.remove_weight_norm()
|
||||
state_dict_g = torch.load(
|
||||
state_dict_g = torch.load( # nosemgrep: trailofbits.python.pickles-in-pytorch.pickles-in-pytorch
|
||||
"%s/GPT_SoVITS/pretrained_models/gsv-v4-pretrained/vocoder.pth" % (now_dir,),
|
||||
map_location="cpu",
|
||||
weights_only=False,
|
||||
weights_only=True,
|
||||
)
|
||||
print("loading vocoder", hifigan_model.load_state_dict(state_dict_g))
|
||||
if is_half == True:
|
||||
@ -475,7 +476,7 @@ hz = 50
|
||||
|
||||
|
||||
def get_gpt_weights(gpt_path):
|
||||
dict_s1 = torch.load(gpt_path, map_location="cpu", weights_only=False)
|
||||
dict_s1 = torch.load(gpt_path, map_location="cpu", weights_only=True) # nosemgrep: trailofbits.python.pickles-in-pytorch.pickles-in-pytorch
|
||||
config = dict_s1["config"]
|
||||
max_sec = config["data"]["max_sec"]
|
||||
t2s_model = Text2SemanticLightningModule(config, "****", is_train=False)
|
||||
@ -1070,7 +1071,7 @@ def get_tts_wav(
|
||||
|
||||
def handle_control(command):
|
||||
if command == "restart":
|
||||
os.execl(g_config.python_exec, g_config.python_exec, *sys.argv)
|
||||
os.execl(sys.executable, sys.executable, *_startup_argv)
|
||||
elif command == "exit":
|
||||
os.kill(os.getpid(), signal.SIGTERM)
|
||||
exit(0)
|
||||
@ -1297,9 +1298,16 @@ change_gpt_sovits_weights(gpt_path=gpt_path, sovits_path=sovits_path)
|
||||
# --------------------------------
|
||||
app = FastAPI()
|
||||
|
||||
_api_key = os.environ.get("API_KEY", "")
|
||||
|
||||
|
||||
def _check_api_key(x_api_key: str = Header(default="")):
|
||||
if _api_key and x_api_key != _api_key:
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
|
||||
|
||||
@app.post("/set_model")
|
||||
async def set_model(request: Request):
|
||||
async def set_model(request: Request, _auth=Depends(_check_api_key)):
|
||||
json_post_raw = await request.json()
|
||||
return change_gpt_sovits_weights(
|
||||
gpt_path=json_post_raw.get("gpt_model_path"), sovits_path=json_post_raw.get("sovits_model_path")
|
||||
@ -1310,18 +1318,19 @@ async def set_model(request: Request):
|
||||
async def set_model(
|
||||
gpt_model_path: str = None,
|
||||
sovits_model_path: str = None,
|
||||
_auth=Depends(_check_api_key),
|
||||
):
|
||||
return change_gpt_sovits_weights(gpt_path=gpt_model_path, sovits_path=sovits_model_path)
|
||||
|
||||
|
||||
@app.post("/control")
|
||||
async def control(request: Request):
|
||||
async def control(request: Request, _auth=Depends(_check_api_key)):
|
||||
json_post_raw = await request.json()
|
||||
return handle_control(json_post_raw.get("command"))
|
||||
|
||||
|
||||
@app.get("/control")
|
||||
async def control(command: str = None):
|
||||
async def control(command: str = None, _auth=Depends(_check_api_key)):
|
||||
return handle_control(command)
|
||||
|
||||
|
||||
|
||||
64
tests/test_invariant_api.py
Normal file
64
tests/test_invariant_api.py
Normal file
@ -0,0 +1,64 @@
|
||||
import pytest
|
||||
import requests
|
||||
import subprocess
|
||||
import time
|
||||
import os
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def api_server():
|
||||
"""Start the actual API server for testing."""
|
||||
# Start the server in a subprocess
|
||||
proc = subprocess.Popen(
|
||||
["python", "api.py"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
env={**os.environ, "PORT": "9999"}
|
||||
)
|
||||
|
||||
# Wait for server to start
|
||||
time.sleep(2)
|
||||
|
||||
yield proc
|
||||
|
||||
# Cleanup
|
||||
proc.terminate()
|
||||
proc.wait()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("payload", [
|
||||
# Missing token (exact exploit case)
|
||||
{"headers": {}, "command": "restart"},
|
||||
# Malformed token
|
||||
{"headers": {"Authorization": "Bearer invalid_token"}, "command": "kill"},
|
||||
# Expired token (simulated)
|
||||
{"headers": {"Authorization": "Bearer eyJ0eXAiOiJKV1QiLCJhbGciOiJIUzI1NiJ9.eyJleHAiOjE2MDAwMDAwMDB9.invalid"}, "command": "status"},
|
||||
# Valid input with no auth (boundary case)
|
||||
{"headers": {}, "command": None},
|
||||
])
|
||||
def test_unauthenticated_requests_rejected(api_server, payload):
|
||||
"""Invariant: Protected endpoints reject unauthenticated requests with 401/403."""
|
||||
|
||||
# Send request to actual endpoint
|
||||
url = "http://localhost:9999/control"
|
||||
|
||||
if payload["command"] is not None:
|
||||
# POST request
|
||||
response = requests.post(
|
||||
url,
|
||||
json={"command": payload["command"]},
|
||||
headers=payload["headers"],
|
||||
timeout=5
|
||||
)
|
||||
else:
|
||||
# GET request
|
||||
response = requests.get(
|
||||
url,
|
||||
params={"command": ""},
|
||||
headers=payload["headers"],
|
||||
timeout=5
|
||||
)
|
||||
|
||||
# Assert authentication failure
|
||||
assert response.status_code in [401, 403], \
|
||||
f"Expected 401/403 for unauthenticated request, got {response.status_code}"
|
||||
Loading…
x
Reference in New Issue
Block a user