From 0b3c679617d073eb15003a0eec5b7ed54546faad Mon Sep 17 00:00:00 2001 From: Daiki Arai Date: Sat, 30 Dec 2023 15:47:32 +0900 Subject: [PATCH] add CORS settings --- README.md | 4 +++- config.py | 3 ++- default_config.yml | 2 ++ server_fastapi.py | 11 +++++++++++ 4 files changed, 18 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index f2ae0fb..6b4531d 100644 --- a/README.md +++ b/README.md @@ -90,9 +90,11 @@ model_assets ### API Server 構築した環境下で`python server_fastapi.py`するとAPIサーバーが起動します。 - API仕様は起動後に`/docs`にて確認ください。 +デフォルトではCORS設定を全てのドメインで許可しています。 +できる限り、`config.yml`の`server.origins`の値を変更し、信頼できるドメインに制限ください(キーを消せばCORS設定を無効にできます)。 + ## Bert-VITS2 v2.1との関係 基本的にはBert-VITS2 v2.1のモデル構造を少し改造しただけです。[事前学習モデル](https://huggingface.co/litagin/Style-Bert-VITS2-1.0-base)も、実質Bert-VITS2 v2.1と同じものを使用しています(不要な重みを削ってsafetensorsに変換したもの)。 diff --git a/config.py b/config.py index c4478e7..302b2d6 100644 --- a/config.py +++ b/config.py @@ -173,12 +173,13 @@ class Webui_config: class Server_config: def __init__( 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.device: str = device self.language: str = language self.limit: int = limit + self.origins: List[str] = origins @classmethod def from_dict(cls, data: Dict[str, any]): diff --git a/default_config.yml b/default_config.yml index 7258ab2..f7fc9c5 100644 --- a/default_config.yml +++ b/default_config.yml @@ -70,3 +70,5 @@ server: device: "cuda" language: "JP" limit: 100 + origins: + - "*" diff --git a/server_fastapi.py b/server_fastapi.py index d6990cd..28e9a40 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -4,6 +4,7 @@ api服务 多版本多模型 fastapi实现 import argparse from fastapi import FastAPI, Query, Request from fastapi.responses import Response, FileResponse +from fastapi.middleware.cors import CORSMiddleware from io import BytesIO from scipy.io import wavfile import uvicorn @@ -68,6 +69,16 @@ if __name__ == "__main__": load_models(model_holder) limit = config.server_config.limit 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 async def _voice(