diff --git a/CHANGELOG.md b/CHANGELOG.md index 189e0083b..293fc5651 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,7 @@ - **sdnq**: simplify pre-quantization saved config - **attention**: refactor settings and improve handling of attention mechanisms - **lora**: separate fuse setting for native-vs-diffuser implementations + - **auth**: strong-enforce auth check on all api endpoints - **Fixes** - hires strength save/load in metadata, thanks @awsr - fix imgi2img initial scale tab, thanks @awsr diff --git a/cli/api-txt2img.js b/cli/api-txt2img.js index 8d0e9f5d1..7b0f6994a 100755 --- a/cli/api-txt2img.js +++ b/cli/api-txt2img.js @@ -30,10 +30,15 @@ async function main() { const headers = new Headers(); const body = JSON.stringify(sd_options); headers.set('Content-Type', 'application/json'); - if (sd_username && sd_password) headers.set({ Authorization: `Basic ${btoa('sd_username:sd_password')}` }); + if (sd_username && sd_password) { + // const credentials = btoa(`${sd_username}:${sd_password}`); + const credentials = Buffer.from(`${sd_username}:${sd_password}`).toString('base64'); + headers.set('Authorization', `Basic ${credentials}`); + } const res = await fetch(`${sd_url}/sdapi/v1/txt2img`, { method, headers, body }); if (res.status !== 200) { - console.log('Error', res.status); + const err = await res.text(); + console.log('Error', res.status, res.statusText, err); } else { const json = await res.json(); console.log('result:', json.info); diff --git a/javascript/login.js b/javascript/login.js index 3f2ab1f64..a9eeb0be9 100644 --- a/javascript/login.js +++ b/javascript/login.js @@ -4,21 +4,21 @@ const loginCSS = ` left: 0; width: 100%; height: 100%; - background: var(--background-fill-primary); - color: var(--body-text-color-subdued); + background: #222; + color: #ddd; font-family: monospace; z-index: 100; `; const loginHTML = ` -
+

Login

- + - +
- +
`; diff --git a/modules/api/api.py b/modules/api/api.py index a260c22ae..ca03f3433 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -4,7 +4,7 @@ from secrets import compare_digest from fastapi import FastAPI, APIRouter, Depends, Request from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException -from modules import errors, shared, postprocessing +from modules import errors, shared from modules.api import models, endpoints, script, helpers, server, generate, process, control, docs, gpu @@ -60,8 +60,8 @@ class Api: self.add_api_route("/sdapi/v1/txt2img", self.generate.post_text2img, methods=["POST"], response_model=models.ResTxt2Img) self.add_api_route("/sdapi/v1/img2img", self.generate.post_img2img, methods=["POST"], response_model=models.ResImg2Img) self.add_api_route("/sdapi/v1/control", self.control.post_control, methods=["POST"], response_model=control.ResControl) - self.add_api_route("/sdapi/v1/extra-single-image", self.extras_single_image_api, methods=["POST"], response_model=models.ResProcessImage) - self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch) + self.add_api_route("/sdapi/v1/extra-single-image", self.process.extras_single_image_api, methods=["POST"], response_model=models.ResProcessImage) + self.add_api_route("/sdapi/v1/extra-batch-images", self.process.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch) self.add_api_route("/sdapi/v1/preprocess", self.process.post_preprocess, methods=["POST"]) self.add_api_route("/sdapi/v1/mask", self.process.post_mask, methods=["POST"]) self.add_api_route("/sdapi/v1/detect", self.process.post_detect, methods=["POST"]) @@ -117,17 +117,22 @@ class Api: from modules.civitai import api_civitai api_civitai.register_api() - - def add_api_route(self, path: str, endpoint, **kwargs): + def add_api_route(self, path: str, fn, **kwargs): + if self.credentials: + deps = list(kwargs.get('dependencies', [])) + deps.append(Depends(self.auth)) + kwargs['dependencies'] = deps if shared.opts.subpath is not None and len(shared.opts.subpath) > 0: - self.app.add_api_route(f'{shared.opts.subpath}{path}', endpoint, **kwargs) - self.app.add_api_route(path, endpoint, **kwargs) + self.app.add_api_route(f'{shared.opts.subpath}{path}', endpoint=fn, **kwargs) + self.app.add_api_route(path, endpoint=fn, **kwargs) def auth(self, credentials: HTTPBasicCredentials = Depends(HTTPBasic())): - # this is only needed for api-only since otherwise auth is handled in gradio/routes.py + if not self.credentials: + return True if credentials.username in self.credentials: if compare_digest(credentials.password, self.credentials[credentials.username]): return True + shared.log.error(f'API authentication: user="{credentials.username}" password="{credentials.password}"') raise HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": "Basic"}) def get_session_start(self, req: Request, agent: Optional[str] = None): @@ -136,27 +141,6 @@ class Api: shared.log.info(f'Browser session: user={user} client={req.client.host} agent={agent}') return {} - def set_upscalers(self, req: dict): - reqDict = vars(req) - reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None) - reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None) - return reqDict - - def extras_single_image_api(self, req: models.ReqProcessImage): - reqDict = self.set_upscalers(req) - reqDict['image'] = helpers.decode_base64_to_image(reqDict['image']) - with self.queue_lock: - result = postprocessing.run_extras(extras_mode=0, image_folder="", input_dir="", output_dir="", save_output=False, **reqDict) - return models.ResProcessImage(image=helpers.encode_pil_to_base64(result[0][0]), html_info=result[1]) - - def extras_batch_images_api(self, req: models.ReqProcessBatch): - reqDict = self.set_upscalers(req) - image_list = reqDict.pop('imageList', []) - image_folder = [helpers.decode_base64_to_image(x.data) for x in image_list] - with self.queue_lock: - result = postprocessing.run_extras(extras_mode=1, image_folder=image_folder, image="", input_dir="", output_dir="", save_output=False, **reqDict) - return models.ResProcessBatch(images=list(map(helpers.encode_pil_to_base64, result[0])), html_info=result[1]) - def launch(self): config = { "listen": shared.cmd_opts.listen, diff --git a/modules/api/process.py b/modules/api/process.py index 3151907a5..c106a18e2 100644 --- a/modules/api/process.py +++ b/modules/api/process.py @@ -4,8 +4,8 @@ from pydantic import BaseModel, Field # pylint: disable=no-name-in-module from fastapi.responses import JSONResponse from fastapi.exceptions import HTTPException from modules.api.helpers import decode_base64_to_image, encode_pil_to_base64 -from modules import errors, shared -from modules.api import models +from modules import errors, shared, postprocessing +from modules.api import models, helpers processor = None # cached instance of processor @@ -175,3 +175,24 @@ class APIProcess(): raise HTTPException(status_code=400, detail="prompt enhancement: invalid type") res = models.ResPromptEnhance(prompt=prompt, seed=seed) return res + + def set_upscalers(self, req: dict): + reqDict = vars(req) + reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None) + reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None) + return reqDict + + def extras_single_image_api(self, req: models.ReqProcessImage): + reqDict = self.set_upscalers(req) + reqDict['image'] = helpers.decode_base64_to_image(reqDict['image']) + with self.queue_lock: + result = postprocessing.run_extras(extras_mode=0, image_folder="", input_dir="", output_dir="", save_output=False, **reqDict) + return models.ResProcessImage(image=helpers.encode_pil_to_base64(result[0][0]), html_info=result[1]) + + def extras_batch_images_api(self, req: models.ReqProcessBatch): + reqDict = self.set_upscalers(req) + image_list = reqDict.pop('imageList', []) + image_folder = [helpers.decode_base64_to_image(x.data) for x in image_list] + with self.queue_lock: + result = postprocessing.run_extras(extras_mode=1, image_folder=image_folder, image="", input_dir="", output_dir="", save_output=False, **reqDict) + return models.ResProcessBatch(images=list(map(helpers.encode_pil_to_base64, result[0])), html_info=result[1])