add authentication & checking of feature flag

This commit is contained in:
salt
2024-04-16 13:48:12 -07:00
parent f948173645
commit e9856fddb6
2 changed files with 49 additions and 4 deletions
+42 -2
View File
@@ -12,7 +12,7 @@ import random
from typing import Dict, Any, Optional
import mimetypes
from fastapi import FastAPI, File, UploadFile, HTTPException, Form
from fastapi import FastAPI, File, UploadFile, HTTPException, Form, Request
from fastapi.responses import FileResponse
import requests
from jsonpath_ng import jsonpath, parse
@@ -32,6 +32,7 @@ SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 60))
STYLE_LOCK = os.environ.get("STYLE_LOCK", None)
PREVIEW_WORKFLOW_NAME = os.environ.get("PREVIEW_WORKFLOW_NAME", None)
MAPPING_NAME = os.environ.get("MAPPING_NAME", "mapping_two.json")
AUTH_URL = os.environ.get("AUTH_URL", "https://api.fakeyou.com/v1/session")
style_to_filename = {
"anime_2_5d": "1_2.5d_anime_model.json",
@@ -77,6 +78,39 @@ CORS_FRAME_ANCESTORS = [
";",
]
def can_use_vst(cookies):
"""Check if the user is authenticated and if the 'video_style_transfer' feature flag is enabled."""
try:
response = requests.get(AUTH_URL, cookies=cookies)
except requests.exceptions.RequestException as e:
print(f"Error during request: {e}")
return False
if response.status_code == 200:
try:
response_json = response.json()
except ValueError:
print("Invalid JSON response")
return False
if response_json.get("success") and response_json.get("logged_in"):
username = response_json["user"]["core_info"]["username"]
print(f"User is logged in: {username}")
feature_flags = response_json["user"].get("maybe_feature_flags", {})
if "video_style_transfer" in feature_flags:
print("VST feature flag is found")
return True
else:
print("VST feature flag not found")
else:
print("User is not logged in or session is invalid")
else:
print(f"Failed to authenticate ({response.status_code}): {response.text}")
return False
def validate_request(request):
try:
request = json.loads(request)
@@ -108,12 +142,18 @@ def terminate_process(proc):
def signal_handler(signum, frame):
print("Signal received, shutting down.")
terminate_process(comfyui_process)
print("ComfyUI server terminated")
exit(0)
@app.post("/preview/")
async def preview_request(request: str = Form(...), input_file: UploadFile = File(...)):
async def preview_request(request_obj: Request, request: str = Form(...), input_file: UploadFile = File(...)):
cookies = request_obj.cookies
if not can_use_vst(cookies):
raise HTTPException(status_code=401, detail="No Video Style Transfer feature flag found")
request = validate_request(request)
if STYLE_LOCK:
if request["style"] != STYLE_LOCK:
raise HTTPException(status_code=400, detail=f"This endpoint is for '{STYLE_LOCK}' style only")
+7 -2
View File
@@ -2,12 +2,17 @@ import json
import requests
url = 'http://localhost:8000/preview/'
# url = 'http://207.189.112.61:31605/preview'
# url = "http://207.189.112.61:31605/preview"
url = "http://localhost:8000/preview"
payload = {"style": "paper_origami", "positive_prompt": "test pos in", "negative_prompt": "test neg in"}
files = {'input_file': open('test_image.jpg', 'rb')}
response = requests.post(url, data={'request': json.dumps(payload)}, files=files) # Use 'request': payload directly
cookies = {
"session": "<put your session cookie here>"
}
response = requests.post(url, data={'request': json.dumps(payload)}, files=files, cookies=cookies)
if response.status_code == 200:
# Success
with open('preview_image.jpeg', 'wb') as f: