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 = `
-
+
`;
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])