1use rand::Rng;
2use serde_json::{json, Value};
3
4#[derive(Debug, Clone)]
24pub struct Txt2ImgRequest {
25 pub positive_prompt: String,
26 pub negative_prompt: String,
27 pub checkpoint: String,
28 pub width: u32,
29 pub height: u32,
30 pub steps: u32,
31 pub cfg_scale: f64,
32 pub sampler: String,
33 pub scheduler: String,
34 pub seed: i64,
35 pub batch_size: u32,
36 pub filename_prefix: String,
37}
38
39impl Txt2ImgRequest {
40 pub fn new(prompt: impl Into<String>, checkpoint: impl Into<String>) -> Self {
43 Self {
44 positive_prompt: prompt.into(),
45 negative_prompt: String::new(),
46 checkpoint: checkpoint.into(),
47 width: 512,
48 height: 768,
49 steps: 25,
50 cfg_scale: 7.5,
51 sampler: "dpmpp_2m".to_string(),
52 scheduler: "karras".to_string(),
53 seed: -1,
54 batch_size: 1,
55 filename_prefix: "ComfyUI".to_string(),
56 }
57 }
58
59 pub fn negative(mut self, prompt: impl Into<String>) -> Self {
61 self.negative_prompt = prompt.into();
62 self
63 }
64
65 pub fn size(mut self, width: u32, height: u32) -> Self {
67 self.width = width;
68 self.height = height;
69 self
70 }
71
72 pub fn steps(mut self, steps: u32) -> Self {
74 self.steps = steps;
75 self
76 }
77
78 pub fn cfg_scale(mut self, cfg: f64) -> Self {
80 self.cfg_scale = cfg;
81 self
82 }
83
84 pub fn sampler(mut self, sampler: impl Into<String>) -> Self {
86 self.sampler = sampler.into();
87 self
88 }
89
90 pub fn scheduler(mut self, scheduler: impl Into<String>) -> Self {
92 self.scheduler = scheduler.into();
93 self
94 }
95
96 pub fn seed(mut self, seed: i64) -> Self {
98 self.seed = seed;
99 self
100 }
101
102 pub fn batch_size(mut self, size: u32) -> Self {
104 self.batch_size = size;
105 self
106 }
107
108 pub fn filename_prefix(mut self, prefix: impl Into<String>) -> Self {
110 self.filename_prefix = prefix.into();
111 self
112 }
113
114 pub fn build(&self) -> (Value, i64) {
119 let seed = if self.seed < 0 {
120 rand::rng().random_range(0..i64::MAX)
121 } else {
122 self.seed
123 };
124
125 let workflow = json!({
126 "1": {
127 "class_type": "CheckpointLoaderSimple",
128 "inputs": {
129 "ckpt_name": self.checkpoint
130 }
131 },
132 "2": {
133 "class_type": "EmptyLatentImage",
134 "inputs": {
135 "width": self.width,
136 "height": self.height,
137 "batch_size": self.batch_size
138 }
139 },
140 "3": {
141 "class_type": "CLIPTextEncode",
142 "inputs": {
143 "text": self.positive_prompt,
144 "clip": ["1", 1]
145 }
146 },
147 "4": {
148 "class_type": "CLIPTextEncode",
149 "inputs": {
150 "text": self.negative_prompt,
151 "clip": ["1", 1]
152 }
153 },
154 "5": {
155 "class_type": "KSampler",
156 "inputs": {
157 "seed": seed,
158 "steps": self.steps,
159 "cfg": self.cfg_scale,
160 "sampler_name": self.sampler,
161 "scheduler": self.scheduler,
162 "denoise": 1.0,
163 "model": ["1", 0],
164 "positive": ["3", 0],
165 "negative": ["4", 0],
166 "latent_image": ["2", 0]
167 }
168 },
169 "6": {
170 "class_type": "VAEDecode",
171 "inputs": {
172 "samples": ["5", 0],
173 "vae": ["1", 2]
174 }
175 },
176 "7": {
177 "class_type": "SaveImage",
178 "inputs": {
179 "filename_prefix": self.filename_prefix,
180 "images": ["6", 0]
181 }
182 }
183 });
184
185 (workflow, seed)
186 }
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 fn make_request() -> Txt2ImgRequest {
194 Txt2ImgRequest::new(
195 "masterpiece, best quality, a cat",
196 "dreamshaper_8.safetensors",
197 )
198 .negative("lowres, blurry")
199 .size(512, 768)
200 .steps(25)
201 .cfg_scale(7.5)
202 .sampler("dpmpp_2m")
203 .scheduler("karras")
204 .seed(12345)
205 }
206
207 #[test]
208 fn test_build_has_all_nodes() {
209 let (workflow, _) = make_request().build();
210 for i in 1..=7 {
211 assert!(workflow.get(i.to_string()).is_some(), "Missing node {}", i);
212 }
213 }
214
215 #[test]
216 fn test_checkpoint_loader() {
217 let (workflow, _) = make_request().build();
218 assert_eq!(workflow["1"]["class_type"], "CheckpointLoaderSimple");
219 assert_eq!(
220 workflow["1"]["inputs"]["ckpt_name"],
221 "dreamshaper_8.safetensors"
222 );
223 }
224
225 #[test]
226 fn test_ksampler_settings() {
227 let (workflow, seed) = make_request().build();
228 let node = &workflow["5"];
229 assert_eq!(node["class_type"], "KSampler");
230 assert_eq!(node["inputs"]["seed"], 12345);
231 assert_eq!(seed, 12345);
232 assert_eq!(node["inputs"]["steps"], 25);
233 assert_eq!(node["inputs"]["cfg"], 7.5);
234 assert_eq!(node["inputs"]["sampler_name"], "dpmpp_2m");
235 assert_eq!(node["inputs"]["scheduler"], "karras");
236 assert_eq!(node["inputs"]["denoise"], 1.0);
237 }
238
239 #[test]
240 fn test_random_seed_when_negative() {
241 let (workflow, seed) = make_request().seed(-1).build();
242 assert!(seed >= 0, "Random seed should be non-negative");
243 assert_eq!(workflow["5"]["inputs"]["seed"], seed);
244 }
245
246 #[test]
247 fn test_clip_text_encode() {
248 let (workflow, _) = make_request().build();
249 assert_eq!(
250 workflow["3"]["inputs"]["text"],
251 "masterpiece, best quality, a cat"
252 );
253 assert_eq!(workflow["3"]["inputs"]["clip"], json!(["1", 1]));
254 assert_eq!(workflow["4"]["inputs"]["text"], "lowres, blurry");
255 }
256
257 #[test]
258 fn test_empty_latent_image() {
259 let (workflow, _) = make_request().build();
260 assert_eq!(workflow["2"]["inputs"]["width"], 512);
261 assert_eq!(workflow["2"]["inputs"]["height"], 768);
262 assert_eq!(workflow["2"]["inputs"]["batch_size"], 1);
263 }
264
265 #[test]
266 fn test_node_connections() {
267 let (workflow, _) = make_request().build();
268 assert_eq!(workflow["5"]["inputs"]["model"], json!(["1", 0]));
269 assert_eq!(workflow["5"]["inputs"]["positive"], json!(["3", 0]));
270 assert_eq!(workflow["5"]["inputs"]["negative"], json!(["4", 0]));
271 assert_eq!(workflow["5"]["inputs"]["latent_image"], json!(["2", 0]));
272 assert_eq!(workflow["6"]["inputs"]["samples"], json!(["5", 0]));
273 assert_eq!(workflow["6"]["inputs"]["vae"], json!(["1", 2]));
274 assert_eq!(workflow["7"]["inputs"]["images"], json!(["6", 0]));
275 }
276
277 #[test]
278 fn test_custom_filename_prefix() {
279 let (workflow, _) = make_request().filename_prefix("MyProject").build();
280 assert_eq!(workflow["7"]["inputs"]["filename_prefix"], "MyProject");
281 }
282
283 #[test]
284 fn test_default_filename_prefix() {
285 let (workflow, _) = Txt2ImgRequest::new("test", "ckpt.safetensors")
286 .seed(1)
287 .build();
288 assert_eq!(workflow["7"]["inputs"]["filename_prefix"], "ComfyUI");
289 }
290
291 #[test]
292 fn test_defaults() {
293 let req = Txt2ImgRequest::new("test prompt", "model.safetensors");
294 assert_eq!(req.width, 512);
295 assert_eq!(req.height, 768);
296 assert_eq!(req.steps, 25);
297 assert_eq!(req.cfg_scale, 7.5);
298 assert_eq!(req.sampler, "dpmpp_2m");
299 assert_eq!(req.scheduler, "karras");
300 assert_eq!(req.seed, -1);
301 assert_eq!(req.batch_size, 1);
302 assert!(req.negative_prompt.is_empty());
303 }
304
305 #[test]
306 fn test_workflow_roundtrip() {
307 let (workflow, _) = make_request().build();
308 let json_str = serde_json::to_string(&workflow).unwrap();
309 let _: Value = serde_json::from_str(&json_str).unwrap();
310 }
311}