From 3f7db39fda64f3c17b46b86b29ad45744d545dea Mon Sep 17 00:00:00 2001 From: drunkplato <6413077+drunkplato@users.noreply.github.com> Date: Tue, 21 Jan 2025 19:00:08 +0000 Subject: [PATCH] logic to upload images from this server --- .gitignore | 3 +- .../__pycache__/autoencoder.cpython-312.pyc | Bin 0 -> 12968 bytes comfy/ldm/models/autoencoder.py | 38 +++ .../MemedeckComfyNodes/nodes_preprocessing.py | 3 + memedeck.py | 260 +++++++++--------- 5 files changed, 178 insertions(+), 126 deletions(-) create mode 100644 comfy/ldm/models/__pycache__/autoencoder.cpython-312.pyc diff --git a/.gitignore b/.gitignore index 626a0b502..189277106 100644 --- a/.gitignore +++ b/.gitignore @@ -14,6 +14,8 @@ __pycache__/ !custom_nodes/example_node.py.example !custom_nodes/MemedeckComfyNodes/ !custom_nodes/MemedeckComfyNodes/** +!comfy/ldm/models/autoencoder.py +!comfy/ldm/models/ extra_model_paths.yaml /.vs @@ -43,4 +45,3 @@ comfy_venv_3.11 models-2 !comfy/ldm/models/autoencoder.py - diff --git a/comfy/ldm/models/__pycache__/autoencoder.cpython-312.pyc b/comfy/ldm/models/__pycache__/autoencoder.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0c7be7c7f191aa4eaa8374f69b2ab531ff5393c8 GIT binary patch literal 12968 zcmdU0Yit|Wm7XWx;!AIdw#Jq$i;gTgwi7#!>R7Sl#I~Z;kEE@dmP>Jl5@kM=87j88 z)HVr>mP(Ua+O#VgZDAQ$tZg`pI=`C!+H|)I1n3`fX+!MP#Tw`qXt6*+M-Ea1`(w|! z!x>UE<))kcv3q5nxpQB0ALpF=opY~$+0bC4AectKJ@ubkDC)QPq6fYlS>H}m)Cwh1 zA|0hB@X18!7&F1p7-yqwjGN$M`~)8}O_*Zl33JRcVToBMtR$a{+G6$zJBjmAN33C@ zfy7NwXUsL>B5^wEp71c#bChVlNQoASn(5IyQF+evR^oc>Cac9Jnn-Oc)VBQ(YB!VG zcBt);nvFJFNNNM5I;9pPwUwm0Ak{6k8mVpbwBPdw?98~IR-K_lJSkmB#)9$SlqCC2 zDmxsXSGkc$D5-K|5hbZwUYt!v67gVE<C=NxQ~6~*36bx>I!e?8 zEm9MVXclRS4Ko=kOpDCLh6zq&B|glE+(mA}B=VAZ#$B#k^?)+g`ZSY*=F zJHVPZ%kMOmrqfiws69{JpvV2J$}4lT(7GiMh{PkwKtOgu1xV)stg9_k_jmuf){yU zc3y6TMp8YG80&!E83&Y9nt}Z3b?UCAVbPqowB5FJ6kB{5?j8F&WKgO}mXdRFJOsjL z1ng!v{MJ!EPEs_AB|Qpa_}7s0cgza&ZJvkJdFlZ5Hvf<8dD_p6-=O3UD1C#Lo8X~f zbAsFgRBfA*l7XZgjKqVJQ7JGRl!GxTDalG||Dy)6SuNxir2)wAaahI1=1UE4HGHPv zSl=|c=?v}U7U8MVsrqXT`UJ+4G{&p1@B)*>ag&gyZFJfH>`AjTy6UPV z8~n0B$K6oET$NRh8sDQK-};DrE@{_uXK-{iA(GTqJd~aQ9yn;5fq1vnp$IA4*5@I&VlksCgjkx$|`d);P~m_ ztmNm(I?JHVDV2>Trc~=}LV;C{BxJb@5*3^r0fWW`=5RtjAC$#Z_oEiQTnJfQi2(t= zb$Y)W&aiiz_T2Imn)YR^MZ4?L({DXpbhod$g}hrRxVJA3-*dOGIG3FTw~!ewdOKJ8 zm;2ux$e8bXTUY$c{_JxFZ_n!aYYSHwZk@{SJdpPu$e4@mLWa#8Uvw5-?xpU`YncYQ z6~5@F&-xosp5wr!fq-fY1YjEGfN=w=BM^9XE*LH6aFdBdlt`nD3?loGpxCIUvr=4v zd9YxLfTRzT8_J3xBt@fv07wdoCXtjQWR7H9MM7}ll_Nm@bdkExF#MKm@)HUkB@4~B zT<>0|;88N!_?B$*72joFiGt|$9k=+OI^K7DAm;Wwo8NWxQ%od7B29*%H6!S3x&F_P zv_V_b6RefOKhXs`m=hV))yPO9N8U{3H(oM@O+=@gK&P8Uvt$9CZUK#K6=8pzut|2& zE!r;9I=fFepiILIPEt*PdUmMidPKeh^4%irkrSTsmzbBT7h)#6Zb`tSJ5DvvcvIW5>TNXllG)kdmPR;d#*H5_}|y{Rsyna&%sp zlH!sajP?t`xF{&IQYaFRgoLmZ1kEfdq^Uk>U}|9FM+Bl%p`Xw)?hYgd`v{5oh1i^u z6ecAY1Z;aT36O#aLlcy_Niujz86XXht5zu%3_!i$ylRdmU>8fssx=ayol6GJN%Ns< z<%jADd2_uDX0~R-SV0-~ot&StActn1=q6pKV zz}};8mbYQ3Ym`qKDsaYARk<7$nI!5I5|Su(B*6FTDApNN12sYC;+y4dL4W=bkm-+E z{SCLodCJcrIOEBTrrakbX%=~D^o3zXNQ8yY=mEb&_Coub>evr4)jX$20aylkAEslU zR4d^cOjLPLud{&bBJpqnbgdLuxhR4=^Wc0=w)3CuMc&@Znzj_6+y7v)C$J=W?z65O);q-kg1F zv2%N706W)AqQdn)RuZPV`M*tpGP^s?P*)V*`2d= zliogrq7V(P9|5vL0kgrgHppH`H$p8kfUyu`O!si|x{YLxe=qFR%PgF$nups}pdwHJ1zYFZw zsPi*sqca)SJS2DKa&kR8pOXgwv7&`x9oNL4k_3oawhzWCJiN z7E)~$TSGO660`FGNInFsAU_R6HN)nL<#d$*MTkAUQ7Pc=s6b_@zEUu975h9tTtWiy4qFPw)eJUUsh;7qO(5r6WpEAiW+{S+?0*v&{;m1Zg=${X8 zM>*=}dxyR7@-N(Q<47y}i&hq1>Z?Q4=+$~#G>ok`(U45-)@27-$uLyXFmLrTx*AN=;L8D% zGJ+;-9R)4&NkjD-$f>OI|DOIH4T#+|o^nJJK?1W6A3iLorX*}+QZm&c3=Itl32q*dDPbw1yKMs*;1?y>(<9vl7HYA?fQl$U@I|H~!YqHXDw)$Y~ltGidba(hnX z8c*J~j1^r2w_RPU?CSGZ8w#$0qQiTiW!$!6YkP@2w+L6ha``K(Uwtom^R>d3CrM_B zVyw1@85QUF^I#)QXgaZX2x!aG(4*c<_45|mV5y^+@|`3kfok0K#p1xxH*85gBf11A6*p_X|&Sg9E zjXO#_rj$&Sw{`KD6G?K)bamADcMdR&ro_p--Vq51B{O@euUc2s%z#-w$h7 zYrKQOjt*81&GoW@&R=v)@Zc<^CQN0+-dx|c3NCRl-r;AHZ1qb>_D#|qWv8s>nzevI z&rdXfo4dYqcazq^rP}~zy%S9Q&FXFT{kr8PN0>#|1+-P`P6g6;{v!8jJ75D&73_tzyDIq9J;Bk0OLr>`NkSktr&F`T#q&l@-P%b<6Rj0PK81LSyisdrK zhK}-x+XJz0E<%Jj15lTcuxJ8j0O;5>7opNTs(iY7J25u=YIi`@c4Oe}YUe;zzF}}f zzOhNoQ<|nRRTlSpjCui}m_Wqm2{~vwNeCRCny@(d{-Rk(h6Gg|?b(Yz0 z%W!urosZD-PeSXPnU?kTdwn-PQ=TkFJ_jwRPTWA@i~`BLB+E5BgB*uKscm(Jy0PAQ zC`Qg3z%P$Lbuqo>+Ld?hD!BS{mVPn}mA-(}k4FvxSOsuBPpJ)aabrOM&>_g$fM}4} z@V9;e2pO%xXpm42RW5rj&B~56S8?a-qejO!7(@ok=yj1%b!_tr?4b+X0-t7Nben_s z6Y6mKe6cbb@dR1OQHbfTmjtHHB6$@F@rr|aLB^*Tz87tuB^y_&z0PcwAuF_?90T$? zRdP@bo=ankV_DC4U(C?IvbPl*o4zl;8(w-U=kOJq+t-@=^38qMJ%#4S7Ht_ibKqC@ zb||tmwR9-w*aDvIj78f;utyv0b({ngkM({aD-^SdrHD`uWSBu2hO-b>p9q>Tj*hOd zbxl-oA~vEcvPIV&t|pEdtlbbRl>Fb5&gW8%G zo$VR%(z(TR*=>2}!Ccn?pc$4NYb>45d@Exl2Nw30K z$Q%TrRh?xR6F^Qi>O(FG>mm6o_$j{xVn8m36HtzcYTr^aW6XjvE5>ZpmTjNdA@;FA z^Sgodg&oDM-S@2y-gD3HzBIHrRN^6Cu8T1X)#fWkGv}ED=uA--9#kcG$u@Jv! zb?WV1dFApe*Z1dKgL&(oPkG4GBnn3tf}asknlP;22C@;CRULpk^aF5mc)|fT15_s3 zHa!S80Ul$&ARaS`W`f5oo8T(8emdYLE8;PNn_M6VcD3zcMjaz41y2zsB{3jIVjIL1 zcTJ*ZE`rWRkdS6(>S%8fP6u=`g-zWku}NtIAS|#ISg^Cby^}PkUmIo+~9M$B|DUzn&;DE6i z-~Zp~2^{E_zb=}qzGB$l3h^s212Wpze?jdLq9V~5Z-A~p;G#oU5_5P^-UY+_bX8{rU++%eF6g5*kuYC4*TzqVIY zZE$oOPa=?gO?BZ}HJ;^`6J%UP{YV?*+Ge(ZF>*)?i%8W%zdCqGlLZ3+{s3wzcYpu@ zEN<<&GIMz*=YQ&h{e`VZGTd!TXR)z;<>`bB2pRwpSG~hJXYP;7lR_xe*pJ#ne zP;Q+s(E8hbJLHzUl&ksUMut{>%Qgjsy9Q18W_J^Bsqa4bG*u?5@0LN4{al z^|o96dH*wc$1@v~+w+dSMMuNBnd=0Y`L^qU4HWK``OEWbT|@bC)40(mI8fdKtUTV$=YcAC+&}r}m=|>y5rfLf! z%Amd*O0srcMDc5<9fZxa6~WIKP<#-KA)U2*mu3ZIG@oounR!vA7_RRttznP}%9g1qV zJK$zV-BXB#DhJ%3IXM7`i+4nH_^Uq`Oi4rS;Y48oW$V#mpz=(5$r0)s!;oj*Y_w@J zd;vz8Hj4~VmgKF&_~L?PlUcCnUw{(q3<`TqNE;<+>TPgAR}OY(@NBRqfSNQ8j>)Ya z)YvtlGV^7uXeDL8Dl@65*0`pvB}Fwyw5V!M zX%QKrrGg3`^^+E~BVNR5sjf|hLNhjgg0+7SguqLOd(FNzZ{M1;_u>(QW7&~C^{xx9 z)WC~t*|kcSFWY!l%**C1y=;RlPxG3mC-3RWJbSmT^G73Bj$b~$8oa*!YPjIrU1%G~ z9KY*sde@X0zS}9Rwih}d`YLd?_VUP-j;WJGe@hXv=^k#C#?DQ9hG~&*mO~so;Bg$y)UFWk;@zT^_sMp7)L3iUWZ}cOyB-eCm#S z2g++N$ZKO8xs7*p)l=|xXO7-;c{AD}hMzv>ch_2`WR+0uYYT*X{O7hz>z)J5V+%-d zv9zl#7pW4{#CNZDe?q~d)X%^@Iu5VewE*fmFm#`7;ys30i1ERFB}}U_59^wnJd`Q* zyZA=k+=S={PyEx+&xSswAgVzp>`DlJM$ml{>sukI+JYpuB7M;^!4g{x+);i;Z7@_j z`k4*aEVhOoC)m`f^;du(27_}x`ZY>!0-aNh)gTu`gKEck*;%Z4QPDa`9o%460HPnLf_MXuxDcu#e^!12 zTB(439Sv5Tvn=ErXb^cO)v{rB3`kytJY^IJLUMa+_Q||`Td`-?wG&rQTu&EzhBC*N zzLa-t)3Nw{$_!ZC-BxTGFY%n<6CY zfM4WO#s=;WxUQ+SNAmsc8bYwa#Nr+QXk?Ncg26fIJl-eMuABm_qGVib8TT*EM}ta7 z)vl7)*l}>UR=(elR$bK_x$}$vodB+%lREI90HQK7klI{H)a^gKr-o*-YSr?<@{-@g zcP1WJO)rqkvf5QvJaZ#nAR={STx3Bt_D78NnPo_8_ON^nUXeEy^q`hlnx=nEIetU! z_;>2i8g=M*)XTr4_Wp+I`!(f$z+uH^Mjjkv?6fV{)c*+uPtv9T20^#|W&i*H literal 0 HcmV?d00001 diff --git a/comfy/ldm/models/autoencoder.py b/comfy/ldm/models/autoencoder.py index e6493155e..02028ce39 100644 --- a/comfy/ldm/models/autoencoder.py +++ b/comfy/ldm/models/autoencoder.py @@ -1,3 +1,4 @@ +<<<<<<< HEAD import logging import math import torch @@ -7,6 +8,15 @@ from typing import Any, Dict, Tuple, Union from comfy.ldm.modules.distributions.distributions import DiagonalGaussianDistribution from comfy.ldm.util import get_obj_from_str, instantiate_from_config +======= +import torch +from contextlib import contextmanager +from typing import Any, Dict, List, Optional, Tuple, Union + +from comfy.ldm.modules.distributions.distributions import DiagonalGaussianDistribution + +from comfy.ldm.util import instantiate_from_config +>>>>>>> 0e1536b4 (logic to upload images from this server) from comfy.ldm.modules.ema import LitEma import comfy.ops @@ -54,7 +64,11 @@ class AbstractAutoencoder(torch.nn.Module): if self.use_ema: self.model_ema = LitEma(self, decay=ema_decay) +<<<<<<< HEAD logging.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.") +======= + logpy.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.") +>>>>>>> 0e1536b4 (logic to upload images from this server) def get_input(self, batch) -> Any: raise NotImplementedError() @@ -70,14 +84,22 @@ class AbstractAutoencoder(torch.nn.Module): self.model_ema.store(self.parameters()) self.model_ema.copy_to(self) if context is not None: +<<<<<<< HEAD logging.info(f"{context}: Switched to EMA weights") +======= + logpy.info(f"{context}: Switched to EMA weights") +>>>>>>> 0e1536b4 (logic to upload images from this server) try: yield None finally: if self.use_ema: self.model_ema.restore(self.parameters()) if context is not None: +<<<<<<< HEAD logging.info(f"{context}: Restored training weights") +======= + logpy.info(f"{context}: Restored training weights") +>>>>>>> 0e1536b4 (logic to upload images from this server) def encode(self, *args, **kwargs) -> torch.Tensor: raise NotImplementedError("encode()-method of abstract base class called") @@ -86,7 +108,11 @@ class AbstractAutoencoder(torch.nn.Module): raise NotImplementedError("decode()-method of abstract base class called") def instantiate_optimizer_from_config(self, params, lr, cfg): +<<<<<<< HEAD logging.info(f"loading >>> {cfg['target']} <<< optimizer from config") +======= + logpy.info(f"loading >>> {cfg['target']} <<< optimizer from config") +>>>>>>> 0e1536b4 (logic to upload images from this server) return get_obj_from_str(cfg["target"])( params, lr=lr, **cfg.get("params", dict()) ) @@ -114,7 +140,11 @@ class AutoencodingEngine(AbstractAutoencoder): self.encoder: torch.nn.Module = instantiate_from_config(encoder_config) self.decoder: torch.nn.Module = instantiate_from_config(decoder_config) +<<<<<<< HEAD self.regularization = instantiate_from_config( +======= + self.regularization: AbstractRegularizer = instantiate_from_config( +>>>>>>> 0e1536b4 (logic to upload images from this server) regularizer_config ) @@ -162,6 +192,7 @@ class AutoencodingEngineLegacy(AutoencodingEngine): }, **kwargs, ) +<<<<<<< HEAD if ddconfig.get("conv3d", False): conv_op = comfy.ops.disable_weight_init.Conv3d @@ -169,12 +200,19 @@ class AutoencodingEngineLegacy(AutoencodingEngine): conv_op = comfy.ops.disable_weight_init.Conv2d self.quant_conv = conv_op( +======= + self.quant_conv = comfy.ops.disable_weight_init.Conv2d( +>>>>>>> 0e1536b4 (logic to upload images from this server) (1 + ddconfig["double_z"]) * ddconfig["z_channels"], (1 + ddconfig["double_z"]) * embed_dim, 1, ) +<<<<<<< HEAD self.post_quant_conv = conv_op(embed_dim, ddconfig["z_channels"], 1) +======= + self.post_quant_conv = comfy.ops.disable_weight_init.Conv2d(embed_dim, ddconfig["z_channels"], 1) +>>>>>>> 0e1536b4 (logic to upload images from this server) self.embed_dim = embed_dim def get_autoencoder_params(self) -> list: diff --git a/custom_nodes/MemedeckComfyNodes/nodes_preprocessing.py b/custom_nodes/MemedeckComfyNodes/nodes_preprocessing.py index 314090456..07a68e31c 100644 --- a/custom_nodes/MemedeckComfyNodes/nodes_preprocessing.py +++ b/custom_nodes/MemedeckComfyNodes/nodes_preprocessing.py @@ -339,6 +339,7 @@ class MD_CompressAdjustNode: image_cv2 = cv2.cvtColor(np.array(tensor2pil(image)), cv2.COLOR_RGB2BGR) # calculate the crf based on the image analysis_results = self.analyze_compression_artifacts(image_cv2, width=width, height=height) + logger.info(f"compression analysis_results: {analysis_results}") calculated_crf = self.calculate_crf(analysis_results, self.ideal_blockiness, self.ideal_edge_density, self.ideal_color_variation, self.blockiness_weight, self.edge_density_weight, self.color_variation_weight) @@ -346,6 +347,8 @@ class MD_CompressAdjustNode: if desired_crf is 0: desired_crf = calculated_crf + logger.info(f"calculated_crf: {calculated_crf}") + # logger.info(f"desired_crf: {desired_crf}") args = [ utils.ffmpeg_path, "-v", "error", diff --git a/memedeck.py b/memedeck.py index 2ad11f16e..8008321f0 100644 --- a/memedeck.py +++ b/memedeck.py @@ -7,6 +7,9 @@ from PIL import Image, ImageOps from functools import partial import pika import json +import numpy as np +from lxml import etree +import io import requests @@ -153,8 +156,10 @@ class MemedeckWorker: routing_key = method.routing_key workflow = 'faceswap' if routing_key == 'faceswap-queue' else 'generation' - user_id = None + user_id = payload["user_id"] if 'user_id' in payload else None + self.logger.info(f"[memedeck]: workflow {workflow} user_id: {user_id}") + if self.video_gen_only: workflow = 'video_gen' user_id = payload["user_id"] @@ -449,7 +454,7 @@ class MemedeckWorker: async def send_preview(self, image_data, sid=None, progress=None, context=None, workflow=None): - self.logger.info(f"[memedeck]: send_preview: {sid}") + # self.logger.info(f"[memedeck]: send_preview: {sid}") if sid is None: self.logger.warning("Received preview without sid") return @@ -483,16 +488,31 @@ class MemedeckWorker: kind = "image_generating" if progress < 100 else "image_generated" + image_id = None + url = None + watermarked_url = None + if kind == "image_generated" and task['workflow'] != 'faceswap': + image_uuid = str(uuid.uuid4()).replace("-", "_") # create uuid for the image + blob_name = f"{task['user_id']}/{image_uuid}" + # upload to azure blob storage + url = await self.azure_storage.save_image(blob_name + ".jpeg", "image/jpeg", preview_bytes) + watermarked_url = await self.azure_storage.save_image_watermarked(blob_name + "_watermarked.jpeg", "image/jpeg", preview_bytes) + image_id = f"image:{image_uuid}" + ai_queue_progress = { "ws_id": sid, "kind": kind, - "data": list(preview_bytes), + "data": list(preview_bytes) if kind == "image_generating" else None, "progress": int(progress), - "context": context + "context": context, + "user_id": task['user_id'], + "image_id": image_id, + "url": url, + "url_watermarked": watermarked_url } - self.logger.info(f"[memedeck]: progress kind: {kind}") - self.logger.info(f"[memedeck]: progress: {progress}") + # self.logger.info(f"[memedeck]: progress kind: {kind}") + # self.logger.info(f"[memedeck]: progress: {progress}") # set the kind to faceswap_generated if workflow is faceswap if workflow == 'faceswap': ai_queue_progress['kind'] = "faceswap_generated" @@ -501,7 +521,7 @@ class MemedeckWorker: await self.send_to_api(ai_queue_progress) if progress == 100 or workflow == 'faceswap': - del self.tasks_by_ws_id[sid] # Remove the task from tasks_by_ws_id + del self.tasks_by_ws_id[sid] # Remove the task from tasks_by_ws_id # self.logger.info(f"[memedeck]: Task {sid} completed") async def send_to_api(self, data): @@ -559,13 +579,17 @@ class MemedeckWorker: # -------------------------------------------------------------------------- # MemedeckAzureStorage # -------------------------------------------------------------------------- -from azure.storage.blob.aio import BlobClient, BlobServiceClient +from azure.storage.blob.aio import BlobServiceClient from azure.storage.blob import ContentSettings from typing import Optional, Tuple import cairosvg -WATERMARK = '' -WATERMARK_SIZE = 40 +WATERMARK = """ + + + +""" +WATERMARK_SIZE = 40 class MemedeckAzureStorage: def __init__(self): @@ -573,7 +597,9 @@ class MemedeckAzureStorage: self.account = os.getenv('STORAGE_ACCOUNT') self.access_key = os.getenv('STORAGE_ACCESS_KEY') self.container = os.getenv('STORAGE_CONTAINER') - self.logger = logging.getLogger(__name__) + logging.getLogger('azure.core.pipeline.policies.http_logging_policy').setLevel(logging.WARNING) + logging.getLogger("azure.storage.common.storageclient").setLevel(logging.WARNING) + self.logger = logging.getLogger('azure.storage.common') if not all([self.account, self.access_key, self.container]): raise EnvironmentError("Missing STORAGE_ACCOUNT, STORAGE_ACCESS_KEY, or STORAGE_CONTAINER environment variables") @@ -607,6 +633,7 @@ class MemedeckAzureStorage: # Upload the blob try: + # prevent logging the request await blob_client.upload_blob( bytes_data, overwrite=True, @@ -621,126 +648,109 @@ class MemedeckAzureStorage: # Construct and return the blob URL blob_url = f"https://media.memedeck.xyz/{self.container}/{blob_name}" return blob_url + + async def save_image_watermarked( + self, + blob_name: str, + content_type: str, + bytes_data: bytes + ) -> str: + image = Image.open(BytesIO(bytes_data)) + watermarked_image = self.add_watermark_to_image(image) + + # convert pil to bytes + img_byte_arr = BytesIO() # Create an in-memory byte stream + watermarked_image.save(img_byte_arr, format=image.format, quality=100, compress_level=1) # Save the image to the in-memory stream + watermarked_image_bytes = img_byte_arr.getvalue() + + return await self.save_image(blob_name, content_type, watermarked_image_bytes) + + def add_watermark_to_image(self, img, background_brightness=None): + """ + Adds a watermark to a single PIL Image. - # async def add_watermark( - # self, - # base_blob_name: str, - # base_image: bytes - # ) -> str: - # """ - # Adds a watermark to the provided image and uploads the watermarked image. + Args: + img: A PIL Image object. - # Args: - # base_blob_name (str): Original blob name of the image. - # base_image (bytes): Image data in bytes. + Returns: + A PIL Image object with the watermark added. + """ - # Returns: - # str: URL of the watermarked image. - # """ - # # Load the input image - # try: - # img = Image.open(BytesIO(base_image)).convert("RGBA") - # except Exception as e: - # raise Exception(f"Failed to load image: {e}") + padding = 12 + x = img.width - WATERMARK_SIZE - padding + y = img.height - WATERMARK_SIZE - padding - # # Calculate position for the watermark (bottom right corner with padding) - # padding = 12 - # x = img.width - WATERMARK_SIZE - padding - # y = img.height - WATERMARK_SIZE - padding + if background_brightness is None: + background_brightness = self.analyze_background_brightness(img, x, y, WATERMARK_SIZE) - # # Analyze background brightness where the watermark will be placed - # background_brightness = self.analyze_background_brightness(img, x, y, WATERMARK_SIZE) - # self.logger.info(f"Background brightness: {background_brightness}") + # Generate watermark image (replace this with your actual watermark generation) + watermark = self.generate_watermark(WATERMARK_SIZE, background_brightness) - # # Render SVG watermark to PNG bytes using cairosvg - # try: - # watermark_png_bytes = cairosvg.svg2png(bytestring=WATERMARK.encode('utf-8'), output_width=WATERMARK_SIZE, output_height=WATERMARK_SIZE) - # watermark = Image.open(BytesIO(watermark_png_bytes)).convert("RGBA") - # except Exception as e: - # raise Exception(f"Failed to render watermark SVG: {e}") + # Overlay the watermark + img.paste(watermark, (x, y), watermark) - # # Determine watermark color based on background brightness - # if background_brightness > 128: - # # Dark watermark for light backgrounds - # watermark_color = (0, 0, 0, int(255 * 0.65)) # Black with 65% opacity - # else: - # # Light watermark for dark backgrounds - # watermark_color = (255, 255, 255, int(255 * 0.65)) # White with 65% opacity - - # # Apply the watermark color by blending - # solid_color = Image.new("RGBA", watermark.size, watermark_color) - # watermark = Image.alpha_composite(watermark, solid_color) - - # # Overlay the watermark onto the original image - # img.paste(watermark, (x, y), watermark) - - # # Save the watermarked image to bytes - # buffer = BytesIO() - # img = img.convert("RGB") # Convert back to RGB for JPEG format - # img.save(buffer, format="JPEG") - # buffer.seek(0) - # jpeg_bytes = buffer.read() - - # # Modify the blob name to include '_watermarked' - # try: - # if "memes/" in base_blob_name: - # base_blob_name_right = base_blob_name.split("memes/", 1)[1] - # else: - # base_blob_name_right = base_blob_name - # base_blob_name_split = base_blob_name_right.rsplit(".", 1) - # base_blob_name_without_extension = base_blob_name_split[0] - # extension = base_blob_name_split[1] - # except Exception as e: - # raise Exception(f"Failed to process blob name: {e}") - - # watermarked_blob_name = f"{base_blob_name_without_extension}_watermarked.{extension}" - - # # Upload the watermarked image - # try: - # watermarked_blob_url = await self.save_image( - # watermarked_blob_name, - # "image/jpeg", - # jpeg_bytes - # ) - # return watermarked_blob_url - # except Exception as e: - # raise Exception(f"Failed to upload watermarked image: {e}") - - # def analyze_background_brightness( - # self, - # img: Image.Image, - # x: int, - # y: int, - # size: int - # ) -> int: - # """ - # Analyzes the brightness of a specific region in the image. - - # Args: - # img (Image.Image): The image to analyze. - # x (int): X-coordinate of the top-left corner of the region. - # y (int): Y-coordinate of the top-left corner of the region. - # size (int): Size of the square region to analyze. - - # Returns: - # int: Average brightness (0-255) of the region. - # """ - # # Crop the specified region - # sub_image = img.crop((x, y, x + size, y + size)).convert("RGB") - - # # Calculate average brightness using the luminance formula - # total_brightness = 0 - # pixel_count = 0 - # for pixel in sub_image.getdata(): - # r, g, b = pixel - # brightness = (r * 299 + g * 587 + b * 114) // 1000 - # total_brightness += brightness - # pixel_count += 1 - - # if pixel_count == 0: - # return 0 - - # average_brightness = total_brightness // pixel_count - # return average_brightness + return img + def analyze_background_brightness(self, img, x, y, size): + """ + Analyzes the average brightness of a region in the image. + + Args: + img: A PIL Image object. + x: The x-coordinate of the top-left corner of the region. + y: The y-coordinate of the top-left corner of the region. + size: The size of the region (square). + + Returns: + The average brightness of the region as an integer. + """ + region = img.crop((x, y, x + size, y + size)) + pixels = np.array(region) + total_brightness = np.sum( + 0.299 * pixels[:, :, 0] + 0.587 * pixels[:, :, 1] + 0.114 * pixels[:, :, 2] + ) / 1000 + print(f"total_brightness: {total_brightness}") + return max(0, min(255, total_brightness)) + + def generate_watermark(self, size, background_brightness): + """ + Generates a watermark image from an SVG string. + + Args: + size: The size of the watermark (square). + background_brightness: The background brightness at the watermark position. + + Returns: + A PIL Image object representing the watermark. + """ + + # Determine watermark color based on background brightness + watermark_color = (0, 0, 0, 165) if background_brightness > 128 else (255, 255, 255, 165) + + # Parse the SVG string + svg_tree = etree.fromstring(WATERMARK) + + # Find the path element and set its fill attribute + path_element = svg_tree.find(".//{http://www.w3.org/2000/svg}path") + if path_element is not None: + r, g, b, a = watermark_color + fill_color = f"rgba({r},{g},{b},{a/255})" # Convert to rgba string + path_element.set("fill", fill_color) + + # Convert the modified SVG tree back to a string + modified_svg = etree.tostring(svg_tree, encoding="unicode") + + # Render the modified SVG to a PNG image with a transparent background + png_data = cairosvg.svg2png( + bytestring=modified_svg, + output_width=size, + output_height=size, + background_color="transparent" + ) + watermark_img = Image.open(BytesIO(png_data)) + + # Convert the watermark to RGBA to handle transparency + watermark_img = watermark_img.convert("RGBA") + + return watermark_img