add CORS settings

This commit is contained in:
Daiki Arai
2023-12-30 15:47:32 +09:00
parent 2b6ec56bbf
commit 0b3c679617
4 changed files with 18 additions and 2 deletions

View File

@@ -90,9 +90,11 @@ model_assets
### API Server ### API Server
構築した環境下で`python server_fastapi.py`するとAPIサーバーが起動します。 構築した環境下で`python server_fastapi.py`するとAPIサーバーが起動します。
API仕様は起動後に`/docs`にて確認ください。 API仕様は起動後に`/docs`にて確認ください。
デフォルトではCORS設定を全てのドメインで許可しています。
できる限り、`config.yml``server.origins`の値を変更し、信頼できるドメインに制限ください(キーを消せばCORS設定を無効にできます)。
## Bert-VITS2 v2.1との関係 ## Bert-VITS2 v2.1との関係
基本的にはBert-VITS2 v2.1のモデル構造を少し改造しただけです。[事前学習モデル](https://huggingface.co/litagin/Style-Bert-VITS2-1.0-base)も、実質Bert-VITS2 v2.1と同じものを使用しています不要な重みを削ってsafetensorsに変換したもの 基本的にはBert-VITS2 v2.1のモデル構造を少し改造しただけです。[事前学習モデル](https://huggingface.co/litagin/Style-Bert-VITS2-1.0-base)も、実質Bert-VITS2 v2.1と同じものを使用しています不要な重みを削ってsafetensorsに変換したもの

View File

@@ -173,12 +173,13 @@ class Webui_config:
class Server_config: class Server_config:
def __init__( def __init__(
self, self,
port: int = 5000, device: str = "cuda", limit: int = 100, language: str = "JP" port: int = 5000, device: str = "cuda", limit: int = 100, language: str = "JP", origins: List[str] = None
): ):
self.port: int = port self.port: int = port
self.device: str = device self.device: str = device
self.language: str = language self.language: str = language
self.limit: int = limit self.limit: int = limit
self.origins: List[str] = origins
@classmethod @classmethod
def from_dict(cls, data: Dict[str, any]): def from_dict(cls, data: Dict[str, any]):

View File

@@ -70,3 +70,5 @@ server:
device: "cuda" device: "cuda"
language: "JP" language: "JP"
limit: 100 limit: 100
origins:
- "*"

View File

@@ -4,6 +4,7 @@ api服务 多版本多模型 fastapi实现
import argparse import argparse
from fastapi import FastAPI, Query, Request from fastapi import FastAPI, Query, Request
from fastapi.responses import Response, FileResponse from fastapi.responses import Response, FileResponse
from fastapi.middleware.cors import CORSMiddleware
from io import BytesIO from io import BytesIO
from scipy.io import wavfile from scipy.io import wavfile
import uvicorn import uvicorn
@@ -68,6 +69,16 @@ if __name__ == "__main__":
load_models(model_holder) load_models(model_holder)
limit = config.server_config.limit limit = config.server_config.limit
app = FastAPI() app = FastAPI()
allow_origins = config.server_config.origins
if allow_origins:
logger.warning(f"CORS allow_origins={config.server_config.origins}. If you don't want, modify config.yml")
app.add_middleware(
CORSMiddleware,
allow_origins=config.server_config.origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
app.logger = logger app.logger = logger
async def _voice( async def _voice(