Skip to main content

comfyui_rs/
workflow.rs

1use rand::Rng;
2use serde_json::{json, Value};
3
4/// Builder for a txt2img ComfyUI workflow.
5///
6/// Constructs a standard 7-node pipeline: CheckpointLoader → CLIP encoders
7/// → KSampler → VAEDecode → SaveImage.
8///
9/// # Example
10/// ```
11/// use comfyui_rs::Txt2ImgRequest;
12///
13/// let (workflow, seed) = Txt2ImgRequest::new("a cat in space", "dreamshaper_8.safetensors")
14///     .negative("lowres, blurry")
15///     .size(512, 768)
16///     .steps(25)
17///     .cfg_scale(7.5)
18///     .build();
19///
20/// assert!(seed >= 0);
21/// assert!(workflow.get("1").is_some()); // CheckpointLoader node
22/// ```
23#[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    /// Create a new request with a prompt and checkpoint. Uses sensible defaults
41    /// for all other parameters (512x768, 25 steps, cfg 7.5, dpmpp_2m/karras).
42    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    /// Set the negative prompt.
60    pub fn negative(mut self, prompt: impl Into<String>) -> Self {
61        self.negative_prompt = prompt.into();
62        self
63    }
64
65    /// Set output dimensions.
66    pub fn size(mut self, width: u32, height: u32) -> Self {
67        self.width = width;
68        self.height = height;
69        self
70    }
71
72    /// Set the number of sampling steps.
73    pub fn steps(mut self, steps: u32) -> Self {
74        self.steps = steps;
75        self
76    }
77
78    /// Set the classifier-free guidance scale.
79    pub fn cfg_scale(mut self, cfg: f64) -> Self {
80        self.cfg_scale = cfg;
81        self
82    }
83
84    /// Set the sampler algorithm (e.g. "euler", "dpmpp_2m", "dpmpp_sde").
85    pub fn sampler(mut self, sampler: impl Into<String>) -> Self {
86        self.sampler = sampler.into();
87        self
88    }
89
90    /// Set the noise scheduler (e.g. "normal", "karras", "exponential").
91    pub fn scheduler(mut self, scheduler: impl Into<String>) -> Self {
92        self.scheduler = scheduler.into();
93        self
94    }
95
96    /// Set a specific seed. Use -1 (the default) for random.
97    pub fn seed(mut self, seed: i64) -> Self {
98        self.seed = seed;
99        self
100    }
101
102    /// Set the batch size (number of images per generation).
103    pub fn batch_size(mut self, size: u32) -> Self {
104        self.batch_size = size;
105        self
106    }
107
108    /// Set the output filename prefix in ComfyUI.
109    pub fn filename_prefix(mut self, prefix: impl Into<String>) -> Self {
110        self.filename_prefix = prefix.into();
111        self
112    }
113
114    /// Build the ComfyUI workflow JSON and resolve the seed.
115    ///
116    /// Returns `(workflow_json, actual_seed)`. When `seed` is -1, a random
117    /// seed is generated and returned so it can be stored with the image.
118    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}