DannyLeo/rmbg_api
0
1from os import getenv2from flask import Flask, render_template_string, jsonify, request, Response3from pathlib import Path4from werkzeug.utils import secure_filename5from re import findall6from rembg import remove, new_session7from requests import get8from urllib.parse import urlparse9 10app = Flask(__name__)11 12models_path = Path(getenv('U2NET_HOME', './models '))13models_path.mkdir(parents=True, exist_ok=True)14models = ["birefnet-cod", "birefnet-dis", "birefnet-general", "birefnet-general-lite", "birefnet-hrsod", "birefnet-massive", "birefnet-portrait", "isnet-anime", "isnet-general-use", "silueta", "u2net", "u2net_cloth_seg", "u2net_human_seg", "u2netp"]15# models = [file.name for file in models_path.iterdir() if file.is_file() and file.suffix == '.onnx']16 17session = new_session(getenv('DEFAULT_MODEL', 'u2netp'))18 19def hex_to_rgba(hex_code):20 hex_code = hex_code.lstrip('#')21 22 if len(hex_code) not in (6, 8):23 return24 # raise ValueError("Invalid hex code, hex code must be 6 to 9 characters long.")25 26 # Extract the red, green, and blue components27 r = int(hex_code[0:2], 16) # Convert the first two characters to decimal28 g = int(hex_code[2:4], 16) # Convert the next two characters to decimal29 b = int(hex_code[4:6], 16) # Convert the last two characters to decimal30 31 if len(hex_code) == 8:32 a = int(hex_code[6:8], 16)33 return (r, g, b, a)34 35 return (r, g, b, 255)36 37 38 39@app.route('/')40def home():41 html_content = '''42 <!DOCTYPE html>43 <html lang="en">44 <head>45 <meta charset="UTF-8">46 <meta name="viewport" content="width=device-width, initial-scale=1.0">47 <title>Welcome to rmbg_api</title>48 <style>49 body {50 font-family: Arial, sans-serif;51 text-align: center;52 padding: 50px;53 }54 h1 {55 color: #333;56 }57 p {58 font-size: 18px;59 }60 </style>61 </head>62 <body>63 <h1>Welcome to my rmbg_api Application!</h1>64 <p>This api is based on the rembg python package</p>65 <p>Please head to our project GitHub for documentation:</p>66 <a href="https://github.com/DannyAkintunde/rmbg_web" target="_blank">Project GitHub</a><br/>67 <small>clone this space to use the api</small>68 </body>69 </html>70 '''71 return render_template_string(html_content)72 73@app.route('/api/rmbg', methods=['GET', 'POST'])74def remove_bg():75 global session76 try:77 apikey = getenv('APIKEY')78 if apikey and apikey != request.headers.get('X-API-KEY'):79 return jsonify({'error': 'Invalid apikey'}), 40180 if request.method == 'POST':81 file = request.files.get('file')82 model = request.form.get('model')83 bg_color = request.form.get('bg_color')84 85 if not file:86 return jsonify({'error': 'no file in request'}), 40087 if model:88 if model not in models:89 return jsonify({'error': f'invalid model name avaliable models are {models}'}), 40090 session = new_session(model)91 92 filename = secure_filename(file.filename)93 mime_type = file.content_type94 95 data = file.stream.read()96 elif request.method == 'GET':97 url = request.args.get('url')98 model = request.args.get('model')99 bg_color = request.args.get('bg_color')100 101 if not url:102 return jsonify({'error': 'Url is required'}), 400103 104 parsed_url = urlparse(url)105 if not parsed_url.scheme or not parsed_url.netloc:106 return jsonify({'error': 'Invalid url passed'}), 400107 108 if model:109 if model not in models:110 return jsonify({'error': f'invalid model name avaliable models are {models}'}), 400111 session = new_session(model)112 113 input_image_response = get(url)114 input_image_response.raise_for_status()115 if input_image_response.status_code == 200:116 data = input_image_response.content117 content_disposition = input_image_response.headers.get('Content-Disposition', '')118 filename_search = findall("filename=(.+)", content_disposition)119 filename = filename_search[0].strip('"').strip("'") if filename_search else url.split("/")[-1]120 mime_type = input_image_response.headers.get('Content-Type', 'application/octet-stream')121 122 bg_color = hex_to_rgba(bg_color) if bg_color else None123 result = remove(data, bgcolor=bg_color, session=session, alpha_matting=True, alpha_matting_foreground_threshold=270, alpha_matting_background_threshold=20, alpha_matting_erode_size=11, post_process_mask=True)124 return Response(response=result, mimetype=mime_type, headers={125 'Content-Disposition': f'inline; filename=removedbg_{filename}'126 })127 except Exception as e:128 print(e)129 return jsonify({'error': f'Server error {str(e)}'}), 500130 131if __name__ == '__main__':132 app.run(host=getenv('HOST', '127.0.0.1'),133 port=getenv('PORT', 8000), debug=(getenv('DEBUG', '').lower() == 'true'))