use std::collections::{HashMap, HashSet};
use anyhow::Result;
use serde::{Deserialize, Serialize};
use crate::expand_prompts::{
build_batch_messages_with_context_for_task, build_remix_messages_with_context_for_task,
build_single_messages_for_task,
};
use crate::{ExpandTask, PromptTransformOperation, RemixDimension};
pub const DISCORD_MAX_VARIATIONS: usize = 5;
pub const EXPANSION_CHUNK_SIZE: usize = 4;
pub const EXPANSION_CHUNK_ATTEMPTS: usize = 3;
pub const MAX_EXPANSION_VARIATIONS: usize = 10_000;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ExpansionAttempt {
pub start: usize,
pub total: usize,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct FamilyOverride {
pub word_limit: Option<u32>,
pub style_notes: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ExpandConfig {
pub model_family: String,
pub task: ExpandTask,
pub operation: PromptTransformOperation,
pub remix_dimensions: Vec<RemixDimension>,
pub variations: usize,
pub temperature: f64,
pub top_p: f64,
pub max_tokens: u32,
pub thinking: bool,
pub system_prompt: Option<String>,
pub batch_prompt: Option<String>,
pub family_overrides: HashMap<String, FamilyOverride>,
pub style: Option<String>,
}
impl Default for ExpandConfig {
fn default() -> Self {
Self {
model_family: "flux".to_string(),
task: ExpandTask::TextToImage,
operation: PromptTransformOperation::Expand,
remix_dimensions: Vec::new(),
variations: 1,
temperature: 0.7,
top_p: 0.9,
max_tokens: 300,
thinking: false,
system_prompt: None,
batch_prompt: None,
family_overrides: HashMap::new(),
style: None,
}
}
}
pub fn resolve_remix_dimensions(
requested: &[RemixDimension],
task: ExpandTask,
style_locked: bool,
) -> Result<Vec<RemixDimension>> {
let allowed: &[RemixDimension] = match task {
ExpandTask::TextToImage | ExpandTask::TextToVideo => &[
RemixDimension::Composition,
RemixDimension::Camera,
RemixDimension::Lighting,
RemixDimension::Setting,
RemixDimension::Mood,
RemixDimension::Movement,
RemixDimension::Style,
],
ExpandTask::ImageToVideo
| ExpandTask::VideoToVideo
| ExpandTask::Retake
| ExpandTask::KeyframeInterpolation
| ExpandTask::AudioDrivenVideo
| ExpandTask::ReferenceToAudioVideo => &[RemixDimension::Movement],
ExpandTask::TextToAudio => &[RemixDimension::Mood, RemixDimension::Movement],
};
let candidates = if requested.is_empty() {
allowed
.iter()
.copied()
.filter(|dimension| !(style_locked && *dimension == RemixDimension::Style))
.collect::<Vec<_>>()
} else {
requested.to_vec()
};
let mut resolved = Vec::new();
for dimension in candidates {
anyhow::ensure!(
allowed.contains(&dimension),
"remix dimension '{dimension}' conflicts with the {task} conditioning authority"
);
anyhow::ensure!(
!(style_locked && dimension == RemixDimension::Style),
"remix dimension 'style' cannot vary while a locked style constraint is set"
);
if !resolved.contains(&dimension) {
resolved.push(dimension);
}
}
anyhow::ensure!(
!resolved.is_empty(),
"remix requires at least one safe dimension"
);
Ok(resolved)
}
pub fn remix_dimensions_for_position(
dimensions: &[RemixDimension],
position: usize,
) -> Vec<RemixDimension> {
dimensions
.get(position.saturating_sub(1) % dimensions.len().max(1))
.copied()
.into_iter()
.collect()
}
#[derive(Debug, Clone)]
pub struct ExpandResult {
pub original: String,
pub expanded: Vec<String>,
}
pub trait PromptExpander: Send + Sync {
fn expand(&self, prompt: &str, config: &ExpandConfig) -> Result<ExpandResult>;
}
pub fn validate_expansion_variation_count(variations: usize) -> Result<()> {
anyhow::ensure!(variations > 0, "variations must be at least 1");
anyhow::ensure!(
variations <= MAX_EXPANSION_VARIATIONS,
"variations exceeds the per-request safety limit of {}",
MAX_EXPANSION_VARIATIONS
);
Ok(())
}
pub fn expand_exact_with<F>(config: &ExpandConfig, mut generate: F) -> Result<Vec<String>>
where
F: FnMut(&ExpandConfig, ExpansionAttempt) -> Result<Vec<String>>,
{
validate_expansion_variation_count(config.variations)?;
let mut expanded = Vec::new();
let mut normalized = HashSet::new();
while expanded.len() < config.variations {
let chunk_target = (config.variations - expanded.len()).min(EXPANSION_CHUNK_SIZE);
let mut chunk = Vec::with_capacity(chunk_target);
for _ in 0..EXPANSION_CHUNK_ATTEMPTS {
let missing = chunk_target - chunk.len();
if missing == 0 {
break;
}
let mut attempt_config = config.clone();
attempt_config.variations = missing;
attempt_config.max_tokens = config.max_tokens.saturating_mul(missing as u32);
let attempt_context = ExpansionAttempt {
start: expanded.len() + chunk.len() + 1,
total: config.variations,
};
let attempt = generate(&attempt_config, attempt_context)?;
if attempt.len() > missing {
anyhow::bail!(
"expansion backend returned {} prompts when exactly {missing} were requested",
attempt.len()
);
}
for prompt in attempt {
let key = normalize_expanded_prompt(&prompt);
if !key.is_empty() && normalized.insert(key) {
chunk.push(prompt);
}
}
}
if chunk.len() != chunk_target {
anyhow::bail!(
"expected exactly {} distinct non-empty prompts, but the expansion backend returned {}",
config.variations,
expanded.len() + chunk.len()
);
}
expanded.extend(chunk);
}
debug_assert_eq!(expanded.len(), config.variations);
Ok(expanded)
}
pub fn validate_expanded_prompts(prompts: &[String], expected: usize) -> Result<()> {
let distinct: HashSet<String> = prompts
.iter()
.map(|prompt| normalize_expanded_prompt(prompt))
.filter(|prompt| !prompt.is_empty())
.collect();
anyhow::ensure!(
prompts.len() == expected && distinct.len() == expected,
"Expected exactly {expected} distinct non-empty prompts, but the host returned {}",
distinct.len()
);
Ok(())
}
fn normalize_expanded_prompt(prompt: &str) -> String {
prompt
.chars()
.flat_map(char::to_lowercase)
.filter(|character| character.is_alphanumeric())
.collect::<String>()
}
#[derive(Debug, Serialize, Deserialize)]
struct ChatMessage {
role: String,
content: String,
}
#[derive(Debug, Serialize)]
struct ChatCompletionRequest {
model: String,
messages: Vec<ChatMessage>,
temperature: f64,
top_p: f64,
max_tokens: u32,
#[serde(skip_serializing_if = "std::ops::Not::not")]
enable_thinking: bool,
}
#[derive(Debug, Deserialize)]
struct ChatCompletionResponse {
choices: Vec<ChatChoice>,
}
#[derive(Debug, Deserialize)]
struct ChatChoice {
message: ChatMessageResponse,
}
#[derive(Debug, Deserialize)]
struct ChatMessageResponse {
content: String,
}
pub struct ApiExpander {
endpoint: String,
model: String,
}
impl ApiExpander {
pub fn new(endpoint: &str, model: &str) -> Self {
let endpoint = endpoint.trim_end_matches('/').to_string();
Self {
endpoint,
model: model.to_string(),
}
}
}
impl PromptExpander for ApiExpander {
fn expand(&self, prompt: &str, config: &ExpandConfig) -> Result<ExpandResult> {
let expanded = expand_exact_with(config, |attempt_config, attempt| {
let family_override = attempt_config
.family_overrides
.get(&attempt_config.model_family);
let messages = if attempt_config.operation == PromptTransformOperation::Remix {
build_remix_messages_with_context_for_task(
prompt,
&attempt_config.model_family,
attempt_config.variations,
attempt_config.task,
Some((attempt.start, attempt.total)),
family_override,
attempt_config.style.as_deref(),
&attempt_config.remix_dimensions,
)
} else if attempt.total > 1 {
build_batch_messages_with_context_for_task(
prompt,
&attempt_config.model_family,
attempt_config.variations,
attempt_config.task,
Some((attempt.start, attempt.total)),
attempt_config.batch_prompt.as_deref(),
family_override,
attempt_config.style.as_deref(),
)
} else {
build_single_messages_for_task(
prompt,
&attempt_config.model_family,
attempt_config.task,
attempt_config.system_prompt.as_deref(),
family_override,
attempt_config.style.as_deref(),
)
};
let chat_messages: Vec<ChatMessage> = messages
.into_iter()
.map(|(role, content)| ChatMessage { role, content })
.collect();
let req_body = ChatCompletionRequest {
model: self.model.clone(),
messages: chat_messages,
temperature: attempt_config.temperature,
top_p: attempt_config.top_p,
max_tokens: attempt_config.max_tokens,
enable_thinking: attempt_config.thinking,
};
let url = format!("{}/v1/chat/completions", self.endpoint);
let body = serde_json::to_string(&req_body)?;
let response_text: String = ureq::post(&url)
.header("Content-Type", "application/json")
.send(body.as_str())
.map_err(|e| anyhow::anyhow!("expand API request failed: {e}"))?
.body_mut()
.read_to_string()
.map_err(|e| anyhow::anyhow!("failed to read expand API response: {e}"))?;
let completion: ChatCompletionResponse = serde_json::from_str(&response_text)
.map_err(|e| anyhow::anyhow!("failed to parse expand API response: {e}"))?;
let content = completion
.choices
.first()
.map(|c| c.message.content.clone())
.filter(|c| !c.trim().is_empty())
.ok_or_else(|| {
anyhow::anyhow!(
"expand API returned empty response (no choices or empty content)"
)
})?;
Ok(
if attempt_config.operation == PromptTransformOperation::Remix || attempt.total > 1
{
parse_variations(&content, attempt_config.variations)
} else {
vec![clean_expanded_prompt(&content)]
},
)
})?;
Ok(ExpandResult {
original: prompt.to_string(),
expanded,
})
}
}
pub fn parse_variations_public(text: &str, expected: usize) -> Vec<String> {
parse_variations(text, expected)
}
pub fn clean_expanded_prompt_public(text: &str) -> String {
clean_expanded_prompt(text)
}
fn parse_variations(text: &str, expected: usize) -> Vec<String> {
let trimmed = text.trim();
if let Ok(arr) = serde_json::from_str::<Vec<String>>(trimmed) {
if !arr.is_empty() {
return arr.into_iter().map(|s| clean_expanded_prompt(&s)).collect();
}
}
if let Some(start) = trimmed.find('[') {
if let Some(end) = trimmed.rfind(']') {
if start < end {
let json_slice = &trimmed[start..=end];
if let Ok(arr) = serde_json::from_str::<Vec<String>>(json_slice) {
if !arr.is_empty() {
return arr.into_iter().map(|s| clean_expanded_prompt(&s)).collect();
}
}
}
}
}
let lines: Vec<String> = trimmed
.lines()
.map(|l| l.trim())
.filter(|l| !l.is_empty())
.map(|l| {
let stripped = l
.trim_start_matches(|c: char| c.is_ascii_digit())
.trim_start_matches(['.', ')', ':', '-'])
.trim_start_matches('"')
.trim_end_matches('"')
.trim();
clean_expanded_prompt(stripped)
})
.filter(|l| !l.is_empty())
.collect();
if lines.len() >= expected {
return lines;
}
let paragraphs: Vec<String> = trimmed
.split("\n\n")
.map(|p| clean_expanded_prompt(p.trim()))
.filter(|p| !p.is_empty())
.collect();
if !paragraphs.is_empty() {
return paragraphs;
}
vec![clean_expanded_prompt(trimmed)]
}
fn clean_expanded_prompt(text: &str) -> String {
let singleton = serde_json::from_str::<Vec<String>>(text.trim())
.ok()
.and_then(|items| (items.len() == 1).then(|| items.into_iter().next().unwrap()));
let trimmed = singleton
.as_deref()
.unwrap_or(text)
.trim()
.trim_matches('"')
.trim_matches('\'')
.trim();
let cleaned = if let Some(end_idx) = trimmed.find("</think>") {
trimmed[end_idx + "</think>".len()..].trim()
} else {
trimmed
};
cleaned.split_whitespace().collect::<Vec<_>>().join(" ")
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ExpandSettings {
#[serde(default)]
pub enabled: bool,
#[serde(default = "default_backend")]
pub backend: String,
#[serde(default = "default_expand_model")]
pub model: String,
#[serde(default = "default_api_model")]
pub api_model: String,
#[serde(default = "default_temperature")]
pub temperature: f64,
#[serde(default = "default_top_p")]
pub top_p: f64,
#[serde(default = "default_max_tokens")]
pub max_tokens: u32,
#[serde(default)]
pub thinking: bool,
#[serde(default)]
pub system_prompt: Option<String>,
#[serde(default)]
pub batch_prompt: Option<String>,
#[serde(default)]
pub families: HashMap<String, FamilyOverride>,
}
fn default_backend() -> String {
"local".to_string()
}
fn default_expand_model() -> String {
"qwen3-expand:q8".to_string()
}
fn default_api_model() -> String {
"qwen2.5:3b".to_string()
}
fn default_temperature() -> f64 {
0.7
}
fn default_top_p() -> f64 {
0.9
}
fn default_max_tokens() -> u32 {
300
}
impl Default for ExpandSettings {
fn default() -> Self {
Self {
enabled: false,
backend: default_backend(),
model: default_expand_model(),
api_model: default_api_model(),
temperature: default_temperature(),
top_p: default_top_p(),
max_tokens: default_max_tokens(),
thinking: false,
system_prompt: None,
batch_prompt: None,
families: HashMap::new(),
}
}
}
impl ExpandSettings {
pub fn with_env_overrides(mut self) -> Self {
if let Ok(v) = std::env::var("MOLD_EXPAND") {
self.enabled = matches!(v.trim().to_lowercase().as_str(), "1" | "true" | "yes");
}
if let Ok(v) = std::env::var("MOLD_EXPAND_BACKEND") {
if !v.is_empty() {
self.backend = v;
}
}
if let Ok(v) = std::env::var("MOLD_EXPAND_MODEL") {
if !v.is_empty() {
if self.is_local() {
self.model = v;
} else {
self.api_model = v;
}
}
}
if let Ok(v) = std::env::var("MOLD_EXPAND_TEMPERATURE") {
if let Ok(t) = v.parse::<f64>() {
self.temperature = t;
}
}
if let Ok(v) = std::env::var("MOLD_EXPAND_THINKING") {
self.thinking = matches!(v.trim().to_lowercase().as_str(), "1" | "true" | "yes");
}
if let Ok(v) = std::env::var("MOLD_EXPAND_SYSTEM_PROMPT") {
if !v.is_empty() {
self.system_prompt = Some(v);
}
}
if let Ok(v) = std::env::var("MOLD_EXPAND_BATCH_PROMPT") {
if !v.is_empty() {
self.batch_prompt = Some(v);
}
}
self
}
pub fn to_expand_config(&self, model_family: &str, variations: usize) -> ExpandConfig {
ExpandConfig {
model_family: model_family.to_string(),
task: ExpandTask::for_family(model_family),
operation: PromptTransformOperation::Expand,
remix_dimensions: Vec::new(),
variations,
temperature: self.temperature,
top_p: self.top_p,
max_tokens: self.max_tokens,
thinking: self.thinking,
system_prompt: self.system_prompt.clone(),
batch_prompt: self.batch_prompt.clone(),
family_overrides: self.families.clone(),
style: None,
}
}
pub fn validate_templates(&self) -> Vec<String> {
let mut warnings = Vec::new();
if let Some(ref tmpl) = self.system_prompt {
for placeholder in ["{WORD_LIMIT}", "{MODEL_NOTES}"] {
if !tmpl.contains(placeholder) {
warnings.push(format!(
"system_prompt is missing placeholder {placeholder} — it won't be substituted"
));
}
}
}
if let Some(ref tmpl) = self.batch_prompt {
for placeholder in ["{N}", "{WORD_LIMIT}", "{MODEL_NOTES}"] {
if !tmpl.contains(placeholder) {
warnings.push(format!(
"batch_prompt is missing placeholder {placeholder} — it won't be substituted"
));
}
}
}
warnings
}
pub fn active_model(&self) -> &str {
if self.is_local() {
&self.model
} else {
&self.api_model
}
}
pub fn create_api_expander(&self) -> Result<Option<ApiExpander>, crate::ModelActivationError> {
crate::require_model_activation(self.active_model(), None)?;
Ok(if self.is_local() {
None
} else {
Some(ApiExpander::new(&self.backend, &self.api_model))
})
}
pub fn is_local(&self) -> bool {
self.backend == "local"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::expand_prompts::build_batch_messages_with_context;
#[test]
fn clean_prompt_strips_quotes() {
assert_eq!(clean_expanded_prompt("\"a cat on mars\""), "a cat on mars");
}
#[test]
fn clean_prompt_strips_single_quotes() {
assert_eq!(clean_expanded_prompt("'a cat on mars'"), "a cat on mars");
}
#[test]
fn clean_prompt_unwraps_singleton_json_array() {
assert_eq!(
clean_expanded_prompt(r#"["a cat on mars"]"#),
"a cat on mars"
);
}
#[test]
fn clean_prompt_strips_thinking() {
let input = "<think>hmm let me think</think>\n\na cat on mars";
assert_eq!(clean_expanded_prompt(input), "a cat on mars");
}
#[test]
fn clean_prompt_strips_multiline_thinking() {
let input = "<think>\nstep 1: analyze\nstep 2: expand\n</think>\n\ndetailed prompt here";
assert_eq!(clean_expanded_prompt(input), "detailed prompt here");
}
#[test]
fn clean_prompt_collapses_whitespace() {
assert_eq!(
clean_expanded_prompt("a cat\n\non mars"),
"a cat on mars"
);
}
#[test]
fn clean_prompt_empty_input() {
assert_eq!(clean_expanded_prompt(""), "");
assert_eq!(clean_expanded_prompt(" "), "");
}
#[test]
fn clean_prompt_only_thinking_block() {
let input = "<think>some reasoning</think>";
assert_eq!(clean_expanded_prompt(input), "");
}
#[test]
fn clean_prompt_preserves_content_without_thinking() {
let input = "a beautiful sunset over the ocean, golden light, dramatic clouds";
assert_eq!(clean_expanded_prompt(input), input);
}
#[test]
fn parse_variations_json_array() {
let input = r#"["a cat", "a dog", "a bird"]"#;
let result = parse_variations(input, 3);
assert_eq!(result, vec!["a cat", "a dog", "a bird"]);
}
#[test]
fn parse_variations_embedded_json() {
let input = "Here are 3 prompts:\n[\"a cat\", \"a dog\", \"a bird\"]";
let result = parse_variations(input, 3);
assert_eq!(result, vec!["a cat", "a dog", "a bird"]);
}
#[test]
fn parse_variations_json_with_thinking() {
let input =
"<think>let me think</think>\n\n[\"expanded cat\", \"expanded dog\", \"expanded bird\"]";
let result = parse_variations(input, 3);
assert_eq!(result.len(), 3);
}
#[test]
fn parse_variations_numbered_list() {
let input = "1. a cat on mars\n2. a dog in space\n3. a bird underwater";
let result = parse_variations(input, 3);
assert_eq!(result.len(), 3);
assert!(result[0].contains("cat"));
assert!(result[1].contains("dog"));
assert!(result[2].contains("bird"));
}
#[test]
fn parse_variations_numbered_with_parens() {
let input = "1) a cat\n2) a dog\n3) a bird";
let result = parse_variations(input, 3);
assert_eq!(result.len(), 3);
assert!(result[0].contains("cat"));
}
#[test]
fn parse_variations_numbered_with_quotes() {
let input = "1. \"a cat on mars\"\n2. \"a dog in space\"";
let result = parse_variations(input, 2);
assert_eq!(result.len(), 2);
assert!(!result[0].starts_with('"'));
assert!(result[0].contains("cat"));
}
#[test]
fn parse_variations_paragraph_fallback() {
let input = "A majestic cat sitting on mars\n\nA playful dog floating in space";
let result = parse_variations(input, 2);
assert_eq!(result.len(), 2);
assert!(result[0].contains("cat"));
assert!(result[1].contains("dog"));
}
#[test]
fn parse_variations_single_text_fallback() {
let input = "just a single prompt with no structure";
let result = parse_variations(input, 3);
assert!(!result.is_empty());
assert!(result[0].contains("single prompt"));
}
#[test]
fn parse_variations_empty_json_array_falls_through() {
let input = "[]";
let result = parse_variations(input, 3);
assert!(!result.is_empty());
}
#[test]
fn parse_variations_cleans_each_item() {
let input = r#"[" a cat ", " a dog "]"#;
let result = parse_variations(input, 2);
assert_eq!(result[0], "a cat");
assert_eq!(result[1], "a dog");
}
#[test]
fn parse_variations_repeated_singleton_json_arrays() {
let input = "[\"a cat\"]\n[\"a dog\"]\n[\"a bird\"]";
let result = parse_variations(input, 3);
assert_eq!(result, vec!["a cat", "a dog", "a bird"]);
}
#[test]
fn exact_expansion_chunks_large_batches_and_scales_token_budget() {
let config = ExpandConfig {
variations: 10,
max_tokens: 300,
..Default::default()
};
let mut attempts = Vec::new();
let expanded = expand_exact_with(&config, |attempt, context| {
attempts.push((
attempt.variations,
attempt.max_tokens,
context.start,
context.total,
));
Ok((0..attempt.variations)
.map(|index| format!("prompt {}", context.start + index))
.collect())
})
.unwrap();
assert_eq!(expanded.len(), 10);
assert_eq!(
attempts,
vec![(4, 1200, 1, 10), (4, 1200, 5, 10), (2, 600, 9, 10)]
);
}
#[test]
fn exact_expansion_retries_only_missing_prompts() {
let config = ExpandConfig {
variations: 8,
..Default::default()
};
let mut requested = Vec::new();
let expanded = expand_exact_with(&config, |attempt, context| {
requested.push(attempt.variations);
let returned = match requested.as_slice() {
[4] => 3,
[4, 1] => {
let messages = build_batch_messages_with_context(
"source",
"flux",
attempt.variations,
Some((context.start, context.total)),
None,
None,
None,
);
assert!(messages[0].1.contains("variations 4 through 4 of 8"));
1
}
_ => attempt.variations,
};
Ok((0..returned)
.map(|index| format!("prompt {}", context.start + index))
.collect())
})
.unwrap();
assert_eq!(expanded.len(), 8);
assert_eq!(requested, vec![4, 1, 4]);
}
#[test]
fn exact_expansion_fails_after_bounded_partial_attempts() {
let config = ExpandConfig {
variations: 8,
..Default::default()
};
let mut attempts = 0;
let error = expand_exact_with(&config, |_, _| {
attempts += 1;
Ok(vec!["only one".to_string()])
})
.unwrap_err();
assert_eq!(attempts, EXPANSION_CHUNK_ATTEMPTS);
assert!(
error
.to_string()
.contains("expected exactly 8 distinct non-empty prompts"),
"{error}"
);
}
#[test]
fn exact_expansion_rejects_zero_variations() {
let config = ExpandConfig {
variations: 0,
..Default::default()
};
let error = expand_exact_with(&config, |_, _| unreachable!()).unwrap_err();
assert!(error.to_string().contains("at least 1"));
}
#[test]
fn exact_expansion_rejects_counts_above_the_safety_limit_without_allocating() {
let config = ExpandConfig {
variations: MAX_EXPANSION_VARIATIONS + 1,
..Default::default()
};
let error = expand_exact_with(&config, |_, _| unreachable!()).unwrap_err();
assert!(error.to_string().contains("safety limit"), "{error}");
}
#[test]
fn exact_expansion_rejects_excess_results_instead_of_truncating() {
let config = ExpandConfig {
variations: 2,
..Default::default()
};
let error = expand_exact_with(&config, |_, _| {
Ok(vec!["one".into(), "two".into(), "three".into()])
})
.unwrap_err();
assert!(error.to_string().contains("exactly 2 were requested"));
}
#[test]
fn exact_expansion_retries_duplicates_from_prior_chunks() {
let config = ExpandConfig {
variations: 6,
..Default::default()
};
let mut attempts = 0;
let expanded = expand_exact_with(&config, |attempt, context| {
attempts += 1;
if context.start == 5 && attempts == 2 {
Ok(vec!["prompt 1".into(), "prompt 5".into()])
} else {
Ok((0..attempt.variations)
.map(|index| format!("prompt {}", context.start + index))
.collect())
}
})
.unwrap();
assert_eq!(expanded.len(), 6);
assert_eq!(expanded[4], "prompt 5");
assert_eq!(expanded[5], "prompt 6");
assert_eq!(attempts, 3);
}
#[test]
fn protocol_validation_rejects_duplicate_or_empty_prompts() {
let duplicate = vec!["A cat".into(), " a cat! ".into()];
let error = validate_expanded_prompts(&duplicate, 2).unwrap_err();
assert!(error.to_string().contains("returned 1"), "{error}");
let empty = vec!["A cat".into(), " ".into()];
let error = validate_expanded_prompts(&empty, 2).unwrap_err();
assert!(error.to_string().contains("returned 1"), "{error}");
}
#[test]
fn expand_settings_defaults() {
let settings = ExpandSettings::default();
assert!(!settings.enabled);
assert_eq!(settings.backend, "local");
assert_eq!(settings.model, "qwen3-expand:q8");
assert_eq!(settings.api_model, "qwen2.5:3b");
assert_eq!(settings.temperature, 0.7);
assert_eq!(settings.top_p, 0.9);
assert_eq!(settings.max_tokens, 300);
assert!(!settings.thinking);
assert!(settings.system_prompt.is_none());
assert!(settings.batch_prompt.is_none());
assert!(settings.families.is_empty());
}
#[test]
fn expand_settings_is_local() {
let settings = ExpandSettings::default();
assert!(settings.is_local());
let api_settings = ExpandSettings {
backend: "http://localhost:11434".to_string(),
..Default::default()
};
assert!(!api_settings.is_local());
}
#[test]
fn expand_settings_create_api_expander_none_for_local() {
let settings = ExpandSettings::default();
assert!(settings.create_api_expander().unwrap().is_none());
}
#[test]
fn expand_settings_create_api_expander_some_for_url() {
let settings = ExpandSettings {
backend: "http://localhost:11434".to_string(),
api_model: "llama3:8b".to_string(),
..Default::default()
};
let expander = settings.create_api_expander().unwrap();
assert!(expander.is_some());
}
#[test]
fn expand_settings_gate_the_selected_local_or_api_model() {
let local_h3 = ExpandSettings {
model: "MiniMax-H3".to_string(),
api_model: "ordinary-inactive-api-model".to_string(),
..Default::default()
};
let error = local_h3
.create_api_expander()
.err()
.expect("the selected local H3 model must be gated");
assert!(error
.to_string()
.contains(crate::MINIMAX_H3_AUTHORIZATION_REQUIRED));
let api_h3 = ExpandSettings {
backend: "http://localhost:11434".to_string(),
model: "ordinary-inactive-local-model".to_string(),
api_model: "MiniMaxAI/MiniMax-H3".to_string(),
..Default::default()
};
assert!(api_h3.create_api_expander().is_err());
let inactive_h3 = ExpandSettings {
backend: "http://localhost:11434".to_string(),
model: "MiniMax-H3".to_string(),
api_model: "llama3:8b".to_string(),
..Default::default()
};
assert!(inactive_h3.create_api_expander().unwrap().is_some());
}
#[test]
fn expand_settings_to_expand_config() {
let settings = ExpandSettings {
temperature: 0.5,
top_p: 0.8,
max_tokens: 200,
thinking: true,
..Default::default()
};
let config = settings.to_expand_config("sdxl", 3);
assert_eq!(config.model_family, "sdxl");
assert_eq!(config.variations, 3);
assert_eq!(config.temperature, 0.5);
assert_eq!(config.top_p, 0.8);
assert_eq!(config.max_tokens, 200);
assert!(config.thinking);
}
#[test]
fn expand_settings_serde_roundtrip() {
let mut families = HashMap::new();
families.insert(
"sd15".to_string(),
FamilyOverride {
word_limit: Some(80),
style_notes: Some("Custom SD1.5 notes.".to_string()),
},
);
let settings = ExpandSettings {
enabled: true,
backend: "http://example.com".to_string(),
model: "qwen3-expand-small:q8".to_string(),
api_model: "gpt-4".to_string(),
temperature: 1.2,
top_p: 0.95,
max_tokens: 500,
thinking: true,
system_prompt: Some("Custom system prompt {WORD_LIMIT} {MODEL_NOTES}".to_string()),
batch_prompt: Some("Custom batch {N} {WORD_LIMIT} {MODEL_NOTES}".to_string()),
families,
};
let toml_str = toml::to_string(&settings).unwrap();
let deserialized: ExpandSettings = toml::from_str(&toml_str).unwrap();
assert_eq!(deserialized.enabled, settings.enabled);
assert_eq!(deserialized.backend, settings.backend);
assert_eq!(deserialized.model, settings.model);
assert_eq!(deserialized.api_model, settings.api_model);
assert_eq!(deserialized.temperature, settings.temperature);
assert_eq!(deserialized.max_tokens, settings.max_tokens);
assert_eq!(deserialized.thinking, settings.thinking);
assert_eq!(deserialized.system_prompt, settings.system_prompt);
assert_eq!(deserialized.batch_prompt, settings.batch_prompt);
assert_eq!(deserialized.families.len(), 1);
let sd15 = deserialized.families.get("sd15").unwrap();
assert_eq!(sd15.word_limit, Some(80));
assert_eq!(sd15.style_notes.as_deref(), Some("Custom SD1.5 notes."));
}
#[test]
fn expand_settings_serde_defaults_on_empty() {
let deserialized: ExpandSettings = toml::from_str("").unwrap();
let defaults = ExpandSettings::default();
assert_eq!(deserialized.enabled, defaults.enabled);
assert_eq!(deserialized.backend, defaults.backend);
assert_eq!(deserialized.model, defaults.model);
assert_eq!(deserialized.temperature, defaults.temperature);
}
#[test]
fn api_expander_strips_trailing_slash() {
let expander = ApiExpander::new("http://localhost:11434/", "qwen2.5:3b");
assert_eq!(expander.endpoint, "http://localhost:11434");
}
#[test]
fn api_expander_no_slash_unchanged() {
let expander = ApiExpander::new("http://localhost:11434", "qwen2.5:3b");
assert_eq!(expander.endpoint, "http://localhost:11434");
}
#[test]
fn expand_config_default() {
let config = ExpandConfig::default();
assert_eq!(config.model_family, "flux");
assert_eq!(config.variations, 1);
assert_eq!(config.temperature, 0.7);
assert_eq!(config.max_tokens, 300);
assert!(!config.thinking);
}
#[test]
fn env_override_model_routes_to_local() {
let settings = ExpandSettings::default();
assert!(settings.is_local());
let mut s = settings;
let v = "qwen3-expand-small:q8".to_string();
if s.is_local() {
s.model = v.clone();
} else {
s.api_model = v.clone();
}
assert_eq!(s.model, "qwen3-expand-small:q8");
assert_eq!(s.api_model, "qwen2.5:3b"); }
#[test]
fn env_override_model_routes_to_api() {
let mut s = ExpandSettings {
backend: "http://localhost:11434".to_string(),
..Default::default()
};
assert!(!s.is_local());
let v = "llama3:70b".to_string();
if s.is_local() {
s.model = v.clone();
} else {
s.api_model = v.clone();
}
assert_eq!(s.api_model, "llama3:70b");
assert_eq!(s.model, "qwen3-expand:q8"); }
#[test]
fn to_expand_config_passes_overrides() {
let mut families = HashMap::new();
families.insert(
"flux".to_string(),
FamilyOverride {
word_limit: Some(200),
style_notes: None,
},
);
let settings = ExpandSettings {
system_prompt: Some("Custom {WORD_LIMIT} {MODEL_NOTES}".to_string()),
batch_prompt: Some("Batch {N} {WORD_LIMIT} {MODEL_NOTES}".to_string()),
families,
..Default::default()
};
let config = settings.to_expand_config("flux", 3);
assert_eq!(
config.system_prompt.as_deref(),
Some("Custom {WORD_LIMIT} {MODEL_NOTES}")
);
assert_eq!(
config.batch_prompt.as_deref(),
Some("Batch {N} {WORD_LIMIT} {MODEL_NOTES}")
);
assert_eq!(config.family_overrides.len(), 1);
assert_eq!(
config.family_overrides.get("flux").unwrap().word_limit,
Some(200)
);
}
#[test]
fn expand_config_default_has_no_overrides() {
let config = ExpandConfig::default();
assert!(config.system_prompt.is_none());
assert!(config.batch_prompt.is_none());
assert!(config.family_overrides.is_empty());
}
#[test]
fn expand_config_default_style_is_none() {
assert!(ExpandConfig::default().style.is_none());
}
#[test]
fn to_expand_config_never_sets_style() {
let settings = ExpandSettings::default();
let config = settings.to_expand_config("flux", 3);
assert!(config.style.is_none());
}
#[test]
fn validate_templates_valid() {
let settings = ExpandSettings {
system_prompt: Some("You are a writer. {WORD_LIMIT} words. {MODEL_NOTES}".to_string()),
batch_prompt: Some(
"Generate {N} prompts. {WORD_LIMIT} words. {MODEL_NOTES}".to_string(),
),
..Default::default()
};
assert!(settings.validate_templates().is_empty());
}
#[test]
fn validate_templates_none_is_valid() {
let settings = ExpandSettings::default();
assert!(settings.validate_templates().is_empty());
}
#[test]
fn validate_templates_missing_word_limit() {
let settings = ExpandSettings {
system_prompt: Some("You are a writer. {MODEL_NOTES}".to_string()),
..Default::default()
};
let errors = settings.validate_templates();
assert_eq!(errors.len(), 1);
assert!(errors[0].contains("{WORD_LIMIT}"));
}
#[test]
fn validate_templates_missing_model_notes() {
let settings = ExpandSettings {
system_prompt: Some("You are a writer. {WORD_LIMIT} words.".to_string()),
..Default::default()
};
let errors = settings.validate_templates();
assert_eq!(errors.len(), 1);
assert!(errors[0].contains("{MODEL_NOTES}"));
}
#[test]
fn validate_templates_batch_missing_n() {
let settings = ExpandSettings {
batch_prompt: Some("Generate prompts. {WORD_LIMIT} {MODEL_NOTES}".to_string()),
..Default::default()
};
let errors = settings.validate_templates();
assert_eq!(errors.len(), 1);
assert!(errors[0].contains("{N}"));
}
#[test]
fn validate_templates_batch_missing_all() {
let settings = ExpandSettings {
batch_prompt: Some("Generate prompts.".to_string()),
..Default::default()
};
let errors = settings.validate_templates();
assert_eq!(errors.len(), 3);
}
#[test]
fn family_override_serde_roundtrip() {
let ov = FamilyOverride {
word_limit: Some(100),
style_notes: Some("Be creative.".to_string()),
};
let json = serde_json::to_string(&ov).unwrap();
let deserialized: FamilyOverride = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.word_limit, Some(100));
assert_eq!(deserialized.style_notes.as_deref(), Some("Be creative."));
}
#[test]
fn family_override_partial_toml() {
let toml_str = "word_limit = 75\n";
let ov: FamilyOverride = toml::from_str(toml_str).unwrap();
assert_eq!(ov.word_limit, Some(75));
assert!(ov.style_notes.is_none());
}
#[test]
fn expand_settings_toml_with_families() {
let toml_str = r#"
enabled = true
system_prompt = "Custom prompt. {WORD_LIMIT} words. {MODEL_NOTES}"
[families.sd15]
word_limit = 40
style_notes = "Short keywords only."
[families.flux]
word_limit = 250
"#;
let settings: ExpandSettings = toml::from_str(toml_str).unwrap();
assert!(settings.enabled);
assert!(settings.system_prompt.is_some());
assert_eq!(settings.families.len(), 2);
let sd15 = settings.families.get("sd15").unwrap();
assert_eq!(sd15.word_limit, Some(40));
assert_eq!(sd15.style_notes.as_deref(), Some("Short keywords only."));
let flux = settings.families.get("flux").unwrap();
assert_eq!(flux.word_limit, Some(250));
assert!(flux.style_notes.is_none());
}
#[test]
fn remix_dimensions_are_task_safe_and_deterministic() {
let text = resolve_remix_dimensions(&[], ExpandTask::TextToImage, false).unwrap();
assert!(text.contains(&RemixDimension::Composition));
assert!(text.contains(&RemixDimension::Style));
assert_eq!(
remix_dimensions_for_position(&text, text.len() + 1),
vec![RemixDimension::Composition]
);
let conditioned = resolve_remix_dimensions(&[], ExpandTask::ImageToVideo, false).unwrap();
assert_eq!(conditioned, vec![RemixDimension::Movement]);
let error = resolve_remix_dimensions(
&[RemixDimension::Composition],
ExpandTask::ImageToVideo,
false,
)
.unwrap_err();
assert!(error.to_string().contains("conditioning authority"));
}
#[test]
fn locked_style_cannot_be_a_remix_dimension() {
let defaults = resolve_remix_dimensions(&[], ExpandTask::TextToImage, true).unwrap();
assert!(!defaults.contains(&RemixDimension::Style));
assert!(
resolve_remix_dimensions(&[RemixDimension::Style], ExpandTask::TextToImage, true)
.is_err()
);
}
}