diff --git a/api.py b/api.py index cc0896a2..16a94080 100644 --- a/api.py +++ b/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) diff --git a/tests/test_invariant_api.py b/tests/test_invariant_api.py new file mode 100644 index 00000000..653c7900 --- /dev/null +++ b/tests/test_invariant_api.py @@ -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}" \ No newline at end of file