ClaudeCharest/skprod
0
1import asyncio2import io3import json4import time5from pathlib import Path6 7import requests8from PIL import Image9 10 11def connect(url, api_key):12 for i in range(200):13 response = requests.post(url+ "/heartbeat", headers={'authorization' : api_key})14 if response.json()['heartbeat'] == 'ready':15 return 16 i+=117 time.sleep(3)18 raise Exception("timeout waiting for GPU to be ready")19 20def get_tokens(url, api_key):21 response = requests.get(url + "/tokens", headers={'authorization' : api_key})22 if response.status_code != 200:23 print(response.status_code)24 print(response.content)25 raise Exception("Error")26 return response.json()['tokens']27 28def get_image(presigned_resp):29 response = requests.get(presigned_resp)30 if response.status_code != 200:31 print(response.status_code)32 print(response.content)33 raise Exception("Error")34 return Image.open(io.BytesIO(response.content))35 36def upload_image(url,api_key, file):37 url = "{}/image".format(url)38 #get presigned url for the destination bucket39 with open(file, 'rb') as f:40 filename = file.split("/")[-1]41 files= {'file': (filename, f)}42 resp = requests.post(url, params={'filename': filename}, headers={'authorization' : api_key})43 if resp.status_code == 200:44 resp = resp.json()45 resp = requests.post(resp['url'], data=resp['fields'], files=files)46 if resp.status_code != 204:47 print(resp.status_code)48 print(resp.content)49 raise Exception("Error")50 else:51 print(resp.status_code)52 print(resp.content)53 return filename54 #upload the image55 56def post_bugreport(url, api_key, bug:str):57 response = requests.post(url + "/bugreport", json={'bug': bug}, headers={'authorization' : api_key})58 59def preprocess(url, api_key, image):60 filename = upload_image(url, api_key, image)61 response = requests.post(url + "/preprocess", params={'image': filename}, headers={'authorization' : api_key})62 if response.status_code == 409:63 return None64 if response.status_code != 200:65 print(response.status_code)66 print(response.content)67 raise Exception("Error")68 prompt_id = response.json()['prompt_id']69 while True:70 response = requests.get(url + "/preprocess",params={'prompt_id': prompt_id}, headers={'authorization' : api_key})71 if response.status_code == 200:72 break73 if response.status_code == 202:74 time.sleep(3)75 continue76 else:77 print(response.status_code)78 print(response.content)79 raise Exception("Error")80 images = []81 for item in response.json()['outputs']:82 images.append(get_image(item))83 return images84 85def sketch_to_image(url, api_key , sketch, control_type, guidance, pos_prompt, neg_prompt, nbimage, mask=None):86 params={'sketch':None, 'control_type': control_type}87 #setting pos/neg prompts88 if pos_prompt != "":89 params['pos_prompt'] = pos_prompt90 if neg_prompt != "":91 params['neg_prompt'] = neg_prompt92 if nbimage != "":93 params['nb_images'] = nbimage94 #setting guidance if provided95 if guidance != None:96 guidance_filename = upload_image(url, api_key, guidance)97 params['guidance'] = guidance_filename98 99 sketch_filename = upload_image(url, api_key, sketch)100 params['sketch'] = sketch_filename101 #Post our prompt to the api endpoint102 response = requests.post(url + "/sketch", params=params, headers={'authorization' : api_key})103 if response.status_code == 409:104 return None105 if response.status_code != 200:106 print(response.status_code)107 print(response.content)108 raise Exception("Error")109 prompt_id = response.json()['prompt_id']110 while True:111 response = requests.get(url + "/sketch", params={"prompt_id":prompt_id}, headers={'authorization' : api_key})112 if response.status_code == 200:113 break114 if response.status_code == 202:115 time.sleep(3)116 continue117 else:118 print(response.status_code)119 print(response.content)120 raise Exception("Error")121 images = []122 for item in response.json()['outputs']:123 images.append(get_image(item))124 return images125 126 127 128 129 