1#![allow(dead_code)]
2use luft_core::contract::backend::{AgentStatus, RunContext};
15use luft_core::contract::event::AgentEvent;
16use luft_core::contract::finding::Finding;
17use luft_core::Scheduler;
18use crate::error::ScriptError;
19use crate::sdk::SdkContext;
20use mlua::{Lua, Table, Value};
21use serde::{Deserialize, Serialize};
22use std::sync::atomic::Ordering;
23use std::sync::Arc;
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
27pub struct ConvergeConfig {
28 pub adversarial: bool,
30 pub vote_threshold: f32,
32 pub max_rounds: u32,
34 pub producers_per_item: u32,
36 pub adversaries_per_finding: u32,
38 pub model: Option<String>,
40}
41
42impl Default for ConvergeConfig {
43 fn default() -> Self {
44 Self {
45 adversarial: true,
46 vote_threshold: 0.7,
47 max_rounds: 3,
48 producers_per_item: 1,
49 adversaries_per_finding: 1,
50 model: None,
51 }
52 }
53}
54
55#[derive(Debug, Clone, Serialize, Deserialize)]
57pub struct ConvergeResult {
58 pub surviving_items: Vec<serde_json::Value>,
60 pub findings: Vec<Finding>,
62 pub rounds: u32,
64 pub converged: bool,
66 pub round_stats: Vec<RoundStats>,
68}
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
72pub struct RoundStats {
73 pub round: u32,
74 pub items_input: usize,
75 pub findings_generated: usize,
76 pub findings_survived: usize,
77 pub findings_refuted: usize,
78 pub approval_rate: f32,
79}
80
81struct ConvergeState {
83 items: Vec<serde_json::Value>,
84 findings: Vec<Finding>,
85 round_stats: Vec<RoundStats>,
86 converged: bool,
87}
88
89pub async fn execute_convergence(
97 items: Vec<serde_json::Value>,
98 producer_prompt: &str,
99 adversary_prompt: &str,
100 config: ConvergeConfig,
101 scheduler: &Arc<Scheduler>,
102 run_ctx: &RunContext,
103) -> Result<ConvergeResult, ScriptError> {
104 if items.is_empty() {
105 tracing::debug!("converge: no items, returning empty");
106 return Ok(ConvergeResult {
107 surviving_items: vec![],
108 findings: vec![],
109 rounds: 0,
110 converged: true,
111 round_stats: vec![],
112 });
113 }
114
115 tracing::info!(
116 n_items = items.len(),
117 max_rounds = config.max_rounds,
118 adversarial = config.adversarial,
119 "converge started"
120 );
121
122 let mut state = ConvergeState {
123 items,
124 findings: vec![],
125 round_stats: vec![],
126 converged: false,
127 };
128
129 for round in 1..=config.max_rounds {
130 let items_count = state.items.len();
131 tracing::info!(round, items_count, "converge round started");
132 let round_input = state.items.clone();
133
134 let (findings, _producer_stats) = generate_findings(
136 &round_input,
137 producer_prompt,
138 config.producers_per_item,
139 config.model.clone(),
140 scheduler,
141 run_ctx,
142 )
143 .await;
144
145 let current_findings = findings.clone();
146 state.findings.extend(current_findings);
147
148 if findings.is_empty() {
149 tracing::info!(round, "converge: no findings generated, ending");
150 break;
151 }
152
153 let (surviving_findings, vote_stats) = if config.adversarial {
155 tracing::debug!(
156 round,
157 n_findings = findings.len(),
158 adversaries = config.adversaries_per_finding,
159 "adversarial verification started"
160 );
161 verify_findings(
162 &findings,
163 adversary_prompt,
164 config.adversaries_per_finding,
165 config.vote_threshold,
166 config.model.clone(),
167 scheduler,
168 run_ctx,
169 )
170 .await
171 } else {
172 (findings.clone(), VoteStats::default())
173 };
174
175 let round_stats = RoundStats {
177 round,
178 items_input: items_count,
179 findings_generated: findings.len(),
180 findings_survived: surviving_findings.len(),
181 findings_refuted: findings.len() - surviving_findings.len(),
182 approval_rate: vote_stats.approval_rate,
183 };
184 tracing::info!(
185 round,
186 generated = findings.len(),
187 survived = surviving_findings.len(),
188 refuted = findings.len() - surviving_findings.len(),
189 approval_rate = vote_stats.approval_rate,
190 "converge round finished"
191 );
192 state.round_stats.push(round_stats);
193
194 if surviving_findings.is_empty() {
196 tracing::info!(round, "converge: all findings refuted, converged");
197 state.converged = true;
198 break;
199 }
200
201 if surviving_findings.len() == findings.len() && round > 1 {
203 tracing::info!(round, "converge: no findings refuted, full convergence");
204 state.converged = true;
205 break;
206 }
207
208 state.items = surviving_findings
210 .into_iter()
211 .map(|f| {
212 serde_json::json!({
213 "kind": f.kind,
214 "severity": format!("{:?}", f.severity).to_lowercase(),
215 "title": f.title,
216 "detail": f.detail,
217 "location": f.location,
218 "evidence": f.evidence,
219 "data": f.data
220 })
221 })
222 .collect();
223 }
224
225 tracing::info!(
226 rounds = state.round_stats.len(),
227 converged = state.converged,
228 surviving = state.items.len(),
229 total_findings = state.findings.len(),
230 "converge finished"
231 );
232 Ok(ConvergeResult {
233 surviving_items: state.items,
234 findings: state.findings,
235 rounds: state.round_stats.len() as u32,
236 converged: state.converged,
237 round_stats: state.round_stats,
238 })
239}
240
241async fn generate_findings(
243 items: &[serde_json::Value],
244 prompt_template: &str,
245 producers_per_item: u32,
246 model: Option<String>,
247 scheduler: &Arc<Scheduler>,
248 run_ctx: &RunContext,
249) -> (Vec<Finding>, ProducerStats) {
250 let mut tasks = Vec::new();
252 for item in items {
253 for _ in 0..producers_per_item {
254 let prompt =
255 prompt_template.replace("{item}", &serde_json::to_string(item).unwrap_or_default());
256 let agent_id = uuid::Uuid::now_v7();
257 let task = luft_core::contract::backend::AgentTask {
258 agent_id,
259 phase_id: 2, prompt,
261 model: model.clone(),
262 description: None,
263 role: Some("producer".to_string()),
264 name: None,
265 agent_seq: 0,
266 allowlist: None,
267 workdir: std::path::PathBuf::from("."),
268 mcp_endpoint: None,
269 timeout: None,
270 output_schema: None,
271 };
272 tasks.push((task, None::<String>));
273 }
274 }
275
276 let results = scheduler.run_parallel(run_ctx.run_id, tasks).await;
278 let agents_run = results.len();
279 tracing::debug!(agents_run, "producer agents completed");
280 let mut all_findings = Vec::new();
281
282 for result in results.into_iter().flatten() {
283 all_findings.extend(result.findings);
284 }
285
286 let stats = ProducerStats {
287 items_processed: items.len(),
288 agents_run,
289 findings_generated: all_findings.len(),
290 };
291
292 (all_findings, stats)
293}
294
295async fn verify_findings(
297 findings: &[Finding],
298 prompt_template: &str,
299 adversaries_per_finding: u32,
300 vote_threshold: f32,
301 model: Option<String>,
302 scheduler: &Arc<Scheduler>,
303 run_ctx: &RunContext,
304) -> (Vec<Finding>, VoteStats) {
305 let mut votes: Vec<(Finding, usize)> = findings.iter().map(|f| (f.clone(), 0)).collect();
306 let total_votes_possible = findings.len() * adversaries_per_finding as usize;
307
308 for finding in findings {
309 let mut approval_count = 0usize;
310
311 for _ in 0..adversaries_per_finding {
312 let prompt = prompt_template.replace(
313 "{finding}",
314 &serde_json::to_string(finding).unwrap_or_default(),
315 );
316 let agent_id = uuid::Uuid::now_v7();
317
318 let task = luft_core::contract::backend::AgentTask {
319 agent_id,
320 phase_id: 2,
321 prompt,
322 model: model.clone(),
323 description: None,
324 role: Some("adversary".to_string()),
325 name: None,
326 agent_seq: 0,
327 allowlist: None,
328 workdir: std::path::PathBuf::from("."),
329 mcp_endpoint: None,
330 timeout: None,
331 output_schema: None,
332 };
333
334 let result = scheduler.run_agent(run_ctx.run_id, task, None).await;
335
336 if let Ok(r) = result {
337 if r.status == AgentStatus::Ok {
338 tracing::trace!(%agent_id, "adversary approved finding");
339 approval_count += 1;
340 } else {
341 tracing::trace!(%agent_id, ?r.status, "adversary rejected finding");
342 }
343 }
344 }
345
346 if let Some((_, count)) = votes.iter_mut().find(|(f, _)| f.title == finding.title) {
347 *count = approval_count;
348 }
349 }
350
351 let approval_rate = if total_votes_possible > 0 {
352 votes.iter().map(|(_, c)| *c).sum::<usize>() as f32 / total_votes_possible as f32
353 } else {
354 0.0
355 };
356
357 let surviving: Vec<Finding> = votes
358 .into_iter()
359 .filter(|(_, count)| {
360 let threshold = (count * 100) / (adversaries_per_finding as usize).max(1);
361 threshold as f32 >= vote_threshold * 100.0
362 })
363 .map(|(f, _)| f)
364 .collect();
365
366 tracing::debug!(
367 surviving = surviving.len(),
368 refuted = findings.len() - surviving.len(),
369 approval_rate,
370 "adversarial voting completed"
371 );
372
373 let stats = VoteStats { approval_rate };
374
375 (surviving, stats)
376}
377
378#[derive(Debug, Default)]
379#[allow(dead_code)] struct ProducerStats {
381 items_processed: usize,
382 agents_run: usize,
383 findings_generated: usize,
384}
385
386#[derive(Debug, Default)]
387struct VoteStats {
388 approval_rate: f32,
389}
390
391fn extract_string(s: mlua::String) -> Option<String> {
397 s.to_str().ok().map(|s| s.to_string())
398}
399
400pub fn register_converge_sdk(lua: &Lua, cx: &SdkContext) -> mlua::Result<()> {
405 let globals = lua.globals();
406 let sched = cx.scheduler.clone();
407 let rc = cx.run_ctx.clone();
408 let handle = cx.handle.clone();
409 let events = cx.events();
410 let run_id = cx.run_id();
411 let phase_counter = cx.phase_counter.clone();
412 let span_counter = cx.span_counter.clone();
413
414 let converge_fn = lua.create_function(move |_lua, (items, options): (Table, Table)| {
415 let scheduler = sched.clone();
416 let run_ctx = rc.clone();
417 let handle = handle.clone();
418 let config = parse_converge_options(&options);
420 tracing::debug!(
421 adversarial = config.adversarial,
422 max_rounds = config.max_rounds,
423 "converge SDK invoked"
424 );
425 let phase_id = phase_counter.load(Ordering::Relaxed);
426 let span_id = span_counter.fetch_add(1, Ordering::Relaxed);
427 let max_rounds = config.max_rounds;
428
429 let producer_prompt = options
430 .get::<mlua::String>("producer_prompt")
431 .ok()
432 .and_then(extract_string)
433 .unwrap_or_else(|| {
434 "Analyze the following item and report any findings: {item}".to_string()
435 });
436
437 let adversary_prompt = options
438 .get::<mlua::String>("adversary_prompt")
439 .ok()
440 .and_then(extract_string)
441 .unwrap_or_else(|| {
442 "Review this finding and determine if it is valid or should be refuted: {finding}"
443 .to_string()
444 });
445
446 let items_vec: Vec<serde_json::Value> = items
448 .sequence_values()
449 .filter_map(|v: mlua::Result<Value>| v.ok())
450 .filter_map(|v| lua_value_to_json(&v).ok())
451 .collect();
452
453 let _ = events.send(AgentEvent::ConvergeStarted {
454 run_id,
455 phase_id,
456 span_id,
457 items: items_vec.len(),
458 max_rounds,
459 });
460 let t0 = std::time::Instant::now();
461
462 let result = handle.block_on(execute_convergence(
464 items_vec,
465 &producer_prompt,
466 &adversary_prompt,
467 config,
468 &scheduler,
469 &run_ctx,
470 ));
471 let elapsed_ms = t0.elapsed().as_millis() as u64;
472
473 match result {
474 Ok(ref res) => {
475 tracing::info!(
476 rounds = res.rounds,
477 converged = res.converged,
478 surviving = res.surviving_items.len(),
479 elapsed_ms,
480 "converge SDK completed"
481 );
482 }
483 Err(ref e) => {
484 tracing::error!(error = %e, elapsed_ms, "converge SDK failed");
485 }
486 }
487 match result {
488 Ok(res) => {
489 let _ = events.send(AgentEvent::ConvergeDone {
490 run_id,
491 phase_id,
492 span_id,
493 rounds: res.rounds,
494 converged: res.converged,
495 surviving: res.surviving_items.len(),
496 result: serde_json::json!({
497 "surviving": res.surviving_items,
498 "rounds": res.rounds,
499 "converged": res.converged,
500 "findings": res.findings,
501 }),
502 elapsed_ms,
503 error: None,
504 });
505 let result_table = _lua.create_table()?;
506 let surviving = _lua.create_table()?;
507 for (i, item) in res.surviving_items.iter().enumerate() {
508 let lua_val = json_to_lua_value(_lua, item.clone())?;
509 surviving.set(i + 1, lua_val)?;
510 }
511 result_table.set("surviving", surviving)?;
512 result_table.set("rounds", res.rounds)?;
513 result_table.set("converged", res.converged)?;
514
515 let findings_table = _lua.create_table()?;
516 for (i, finding) in res.findings.iter().enumerate() {
517 let ft = _lua.create_table()?;
518 ft.set("kind", finding.kind.as_str())?;
519 ft.set("severity", format!("{:?}", finding.severity).to_lowercase())?;
520 ft.set("title", finding.title.as_str())?;
521 ft.set("detail", finding.detail.as_str())?;
522 findings_table.set(i + 1, ft)?;
523 }
524 result_table.set("findings", findings_table)?;
525
526 Ok(result_table)
527 }
528 Err(e) => {
529 let _ = events.send(AgentEvent::ConvergeDone {
530 run_id,
531 phase_id,
532 span_id,
533 rounds: 0,
534 converged: false,
535 surviving: 0,
536 result: serde_json::Value::Null,
537 elapsed_ms,
538 error: Some(e.to_string()),
539 });
540 Err(mlua::Error::RuntimeError(format!("converge error: {}", e)))
541 }
542 }
543 })?;
544
545 globals.set("converge", converge_fn)?;
546 Ok(())
547}
548
549fn parse_converge_options(options: &Table) -> ConvergeConfig {
551 let mut config = ConvergeConfig::default();
552
553 if let Ok(adversarial) = options.get::<mlua::String>("adversarial") {
554 let s = extract_string(adversarial).unwrap_or_else(|| "true".to_string());
555 config.adversarial = s != "false";
556 }
557 if let Ok(threshold) = options.get::<mlua::Number>("vote_threshold") {
558 config.vote_threshold = threshold as f32;
559 }
560 if let Ok(max_rounds) = options.get::<mlua::Integer>("max_rounds") {
561 config.max_rounds = max_rounds as u32;
562 }
563 if let Ok(producers) = options.get::<mlua::Integer>("producers") {
564 config.producers_per_item = producers as u32;
565 }
566 if let Ok(adversaries) = options.get::<mlua::Integer>("adversaries") {
567 config.adversaries_per_finding = adversaries as u32;
568 }
569 if let Ok(model) = options.get::<mlua::String>("model") {
570 let s = extract_string(model).unwrap_or_default();
571 config.model = if s.is_empty() { None } else { Some(s) };
572 }
573
574 config
575}
576
577fn lua_value_to_json(value: &Value) -> Result<serde_json::Value, mlua::Error> {
579 match value {
580 Value::Nil => Ok(serde_json::Value::Null),
581 Value::Boolean(b) => Ok(serde_json::Value::Bool(*b)),
582 Value::Integer(i) => Ok(serde_json::Value::Number(serde_json::Number::from(*i))),
583 Value::Number(n) => {
584 if let Some(n) = serde_json::Number::from_f64(*n) {
585 Ok(serde_json::Value::Number(n))
586 } else {
587 Ok(serde_json::Value::Null)
588 }
589 }
590 Value::String(s) => {
591 let owned = s.clone();
592 let s = match owned.to_str() {
593 Ok(s) => s.to_string(),
594 Err(_) => return Ok(serde_json::Value::Null),
595 };
596 Ok(serde_json::Value::String(s))
597 }
598 Value::Table(t) => {
599 let len = t.len().unwrap_or(0);
601 if len > 0 {
602 let arr: Vec<serde_json::Value> = t
603 .sequence_values()
604 .filter_map(|v| v.ok())
605 .filter_map(|v| lua_value_to_json(&v).ok())
606 .collect();
607 if !arr.is_empty() {
608 return Ok(serde_json::Value::Array(arr));
609 }
610 }
611
612 let mut map = serde_json::Map::new();
614 for (k, v) in t.pairs::<Value, Value>().flatten() {
615 let key = match k {
616 Value::String(s) => s
617 .clone()
618 .to_str()
619 .map(|s| s.to_string())
620 .unwrap_or_default(),
621 Value::Integer(i) => i.to_string(),
622 _ => continue,
623 };
624 if let Ok(v) = lua_value_to_json(&v) {
625 map.insert(key, v);
626 }
627 }
628 Ok(serde_json::Value::Object(map))
629 }
630 _ => Ok(serde_json::Value::Null),
631 }
632}
633
634fn json_to_lua_value(lua: &Lua, json: serde_json::Value) -> Result<Value, mlua::Error> {
636 match json {
637 serde_json::Value::Null => Ok(Value::Nil),
638 serde_json::Value::Bool(b) => Ok(Value::Boolean(b)),
639 serde_json::Value::Number(n) => {
640 if let Some(i) = n.as_i64() {
641 Ok(Value::Integer(i))
642 } else if let Some(f) = n.as_f64() {
643 Ok(Value::Number(f))
644 } else {
645 Ok(Value::Nil)
646 }
647 }
648 serde_json::Value::String(s) => Ok(Value::String(lua.create_string(&s)?)),
649 serde_json::Value::Array(arr) => {
650 let t = lua.create_table()?;
651 for (i, v) in arr.into_iter().enumerate() {
652 t.set(i + 1, json_to_lua_value(lua, v)?)?;
653 }
654 Ok(Value::Table(t))
655 }
656 serde_json::Value::Object(map) => {
657 let t = lua.create_table()?;
658 for (k, v) in map {
659 t.set(k, json_to_lua_value(lua, v)?)?;
660 }
661 Ok(Value::Table(t))
662 }
663 }
664}
665
666#[cfg(test)]
667mod tests {
668 use super::*;
669 use luft_core::contract::backend::{
670 AgentBackend, AgentCapabilities, AgentResult, BackendError, LogRef,
671 };
672 use luft_core::contract::finding::{Location, Severity};
673 use luft_core::contract::ids::TokenUsage;
674 use luft_core::scheduler::{BackendRegistry, RetryPolicy, SchedulerConfig};
675 use luft_core::{AgentTask, MockBackend, MockBehavior};
676 use crate::sdk::ReportSink;
677 use async_trait::async_trait;
678 use std::sync::Arc;
679 use std::time::Duration;
680 use tokio::sync::broadcast;
681 use tokio_util::sync::CancellationToken;
682 use uuid::Uuid;
683
684 fn sample_finding(title: &str) -> Finding {
687 Finding {
688 kind: "test".into(),
689 severity: Severity::Info,
690 title: title.into(),
691 detail: "A detailed description".into(),
692 location: Some(Location {
693 file: "src/main.rs".into(),
694 line: Some(42),
695 }),
696 evidence: vec!["line 42: suspect code".into()],
697 data: serde_json::json!({"extra": "info"}),
698 }
699 }
700
701 fn default_finding() -> Finding {
702 sample_finding("Test Finding")
703 }
704
705 fn converge_scheduler(backend: Arc<dyn AgentBackend>) -> Arc<Scheduler> {
708 let config = SchedulerConfig {
709 max_concurrency: 4,
710 quota_per_run: 1000,
711 retry: RetryPolicy::default(),
712 };
713 let registry = BackendRegistry::new().with(backend);
714 Scheduler::new(config, registry, None)
715 }
716
717 fn test_run_ctx(scheduler: &Arc<Scheduler>) -> (RunContext, broadcast::Receiver<AgentEvent>) {
718 let run_id = Uuid::now_v7();
719 let rx = scheduler.init_run(run_id, 64);
720 let (tx, _rx2) = broadcast::channel(64);
721 let ctx = RunContext {
722 run_id,
723 cancel: CancellationToken::new(),
724 events: tx,
725 };
726 (ctx, rx)
727 }
728
729 #[test]
734 fn test_default_config() {
735 let c = ConvergeConfig::default();
736 assert!(c.adversarial);
737 assert!((c.vote_threshold - 0.7).abs() < f32::EPSILON);
738 assert_eq!(c.max_rounds, 3);
739 assert_eq!(c.producers_per_item, 1);
740 assert_eq!(c.adversaries_per_finding, 1);
741 assert!(c.model.is_none());
742 }
743
744 #[test]
745 fn test_converge_config_debug_clone_serialize() {
746 let c = ConvergeConfig::default();
747 let _ = format!("{:?}", c);
748 let _ = c.clone();
749 let json = serde_json::to_string(&c).unwrap();
750 let back: ConvergeConfig = serde_json::from_str(&json).unwrap();
751 assert_eq!(c.max_rounds, back.max_rounds);
752 }
753
754 #[test]
759 fn test_parse_options_defaults() {
760 let lua = Lua::new();
761 let t = lua.create_table().unwrap();
762 let cfg = parse_converge_options(&t);
763 assert!(cfg.adversarial);
764 assert!((cfg.vote_threshold - 0.7).abs() < f32::EPSILON);
765 assert_eq!(cfg.max_rounds, 3);
766 assert_eq!(cfg.producers_per_item, 1);
767 assert_eq!(cfg.adversaries_per_finding, 1);
768 assert!(cfg.model.is_none());
769 }
770
771 #[test]
772 fn test_parse_options_all_fields() {
773 let lua = Lua::new();
774 let t = lua.create_table().unwrap();
775 t.set("adversarial", "false").unwrap();
776 t.set("vote_threshold", 0.5).unwrap();
777 t.set("max_rounds", 5u32).unwrap();
778 t.set("producers", 2u32).unwrap();
779 t.set("adversaries", 3u32).unwrap();
780 t.set("model", "gpt-4").unwrap();
781 let cfg = parse_converge_options(&t);
782 assert!(!cfg.adversarial);
783 assert!((cfg.vote_threshold - 0.5).abs() < f32::EPSILON);
784 assert_eq!(cfg.max_rounds, 5);
785 assert_eq!(cfg.producers_per_item, 2);
786 assert_eq!(cfg.adversaries_per_finding, 3);
787 assert_eq!(cfg.model.as_deref(), Some("gpt-4"));
788 }
789
790 #[test]
791 fn test_parse_options_adversarial_true_string() {
792 let lua = Lua::new();
793 let t = lua.create_table().unwrap();
794 t.set("adversarial", "true").unwrap();
795 assert!(parse_converge_options(&t).adversarial);
796 }
797
798 #[test]
799 fn test_parse_options_empty_model() {
800 let lua = Lua::new();
801 let t = lua.create_table().unwrap();
802 t.set("model", "").unwrap();
803 assert!(parse_converge_options(&t).model.is_none());
804 }
805
806 #[test]
807 fn test_parse_options_model_some() {
808 let lua = Lua::new();
809 let t = lua.create_table().unwrap();
810 t.set("model", "claude").unwrap();
811 assert_eq!(parse_converge_options(&t).model.as_deref(), Some("claude"));
812 }
813
814 #[test]
819 fn test_extract_string_valid() {
820 let lua = Lua::new();
821 assert_eq!(
822 extract_string(lua.create_string("hello").unwrap()),
823 Some("hello".into())
824 );
825 }
826
827 #[test]
828 fn test_extract_string_empty() {
829 let lua = Lua::new();
830 assert_eq!(
831 extract_string(lua.create_string("").unwrap()),
832 Some("".into())
833 );
834 }
835
836 #[test]
837 fn test_extract_string_invalid_utf8() {
838 let lua = Lua::new();
839 let s = lua.create_string([0xFF, 0xFE, 0x00]).unwrap();
840 assert_eq!(extract_string(s), None);
841 }
842
843 #[test]
848 fn test_lua_value_to_json_nil() {
849 assert_eq!(
850 lua_value_to_json(&Value::Nil).unwrap(),
851 serde_json::Value::Null
852 );
853 }
854
855 #[test]
856 fn test_lua_value_to_json_boolean() {
857 assert_eq!(
858 lua_value_to_json(&Value::Boolean(true)).unwrap(),
859 serde_json::Value::Bool(true)
860 );
861 assert_eq!(
862 lua_value_to_json(&Value::Boolean(false)).unwrap(),
863 serde_json::Value::Bool(false)
864 );
865 }
866
867 #[test]
868 fn test_lua_value_to_json_integer() {
869 assert_eq!(
870 lua_value_to_json(&Value::Integer(42)).unwrap(),
871 serde_json::json!(42)
872 );
873 assert_eq!(
874 lua_value_to_json(&Value::Integer(-5)).unwrap(),
875 serde_json::json!(-5)
876 );
877 }
878
879 #[test]
880 fn test_lua_value_to_json_number() {
881 assert_eq!(
882 lua_value_to_json(&Value::Number(std::f64::consts::PI)).unwrap(),
883 serde_json::json!(std::f64::consts::PI)
884 );
885 }
886
887 #[test]
888 fn test_lua_value_to_json_number_nan() {
889 assert_eq!(
890 lua_value_to_json(&Value::Number(f64::NAN)).unwrap(),
891 serde_json::Value::Null
892 );
893 }
894
895 #[test]
896 fn test_lua_value_to_json_number_infinity() {
897 assert_eq!(
898 lua_value_to_json(&Value::Number(f64::INFINITY)).unwrap(),
899 serde_json::Value::Null
900 );
901 assert_eq!(
902 lua_value_to_json(&Value::Number(f64::NEG_INFINITY)).unwrap(),
903 serde_json::Value::Null
904 );
905 }
906
907 #[test]
908 fn test_lua_value_to_json_string_valid() {
909 let lua = Lua::new();
910 let s = lua.create_string("hello world").unwrap();
911 assert_eq!(
912 lua_value_to_json(&Value::String(s)).unwrap(),
913 serde_json::json!("hello world")
914 );
915 }
916
917 #[test]
918 fn test_lua_value_to_json_string_invalid_utf8() {
919 let lua = Lua::new();
920 let s = lua.create_string([0xFF, 0xFE]).unwrap();
921 assert_eq!(
922 lua_value_to_json(&Value::String(s)).unwrap(),
923 serde_json::Value::Null
924 );
925 }
926
927 #[test]
928 fn test_lua_value_to_json_table_as_array() {
929 let lua = Lua::new();
930 let t = lua.create_table().unwrap();
931 t.set(1, "a").unwrap();
932 t.set(2, "b").unwrap();
933 t.set(3, "c").unwrap();
934 assert_eq!(
935 lua_value_to_json(&Value::Table(t)).unwrap(),
936 serde_json::json!(["a", "b", "c"])
937 );
938 }
939
940 #[test]
941 fn test_lua_value_to_json_table_as_object() {
942 let lua = Lua::new();
943 let t = lua.create_table().unwrap();
944 t.set("name", "test").unwrap();
945 t.set("count", 42).unwrap();
946 let val = lua_value_to_json(&Value::Table(t)).unwrap();
947 let obj = val.as_object().unwrap();
948 assert_eq!(obj["name"], "test");
949 assert_eq!(obj["count"], 42);
950 }
951
952 #[test]
953 fn test_lua_value_to_json_table_empty() {
954 let lua = Lua::new();
955 let t = lua.create_table().unwrap();
956 assert_eq!(
957 lua_value_to_json(&Value::Table(t)).unwrap(),
958 serde_json::json!({})
959 );
960 }
961
962 #[test]
963 fn test_lua_value_to_json_table_nested() {
964 let lua = Lua::new();
965 let inner = lua.create_table().unwrap();
966 inner.set("key", "val").unwrap();
967 let outer = lua.create_table().unwrap();
968 outer.set("nested", inner).unwrap();
969 let val = lua_value_to_json(&Value::Table(outer)).unwrap();
970 assert_eq!(val, serde_json::json!({"nested": {"key": "val"}}));
971 }
972
973 #[test]
974 fn test_lua_value_to_json_table_integer_key_in_object() {
975 let lua = Lua::new();
976 let t = lua.create_table().unwrap();
977 t.set(42, "answer").unwrap();
978 let val = lua_value_to_json(&Value::Table(t)).unwrap();
979 assert_eq!(val, serde_json::json!({"42": "answer"}));
980 }
981
982 #[test]
983 fn test_lua_value_to_json_function_catch_all() {
984 let lua = Lua::new();
985 let f = lua.create_function(|_, ()| Ok(())).unwrap();
986 assert_eq!(
987 lua_value_to_json(&Value::Function(f)).unwrap(),
988 serde_json::Value::Null
989 );
990 }
991
992 #[test]
993 fn test_lua_value_to_json_userdata_catch_all() {
994 let lua = Lua::new();
995 let thread = lua
997 .create_thread(lua.load("return 1").into_function().unwrap())
998 .unwrap();
999 let val = Value::Thread(thread);
1000 assert_eq!(lua_value_to_json(&val).unwrap(), serde_json::Value::Null);
1001 }
1002
1003 #[test]
1008 fn test_json_to_lua_value_null() {
1009 let lua = Lua::new();
1010 assert!(matches!(
1011 json_to_lua_value(&lua, serde_json::Value::Null).unwrap(),
1012 Value::Nil
1013 ));
1014 }
1015
1016 #[test]
1017 fn test_json_to_lua_value_bool() {
1018 let lua = Lua::new();
1019 assert!(matches!(
1020 json_to_lua_value(&lua, serde_json::Value::Bool(true)).unwrap(),
1021 Value::Boolean(true)
1022 ));
1023 }
1024
1025 #[test]
1026 fn test_json_to_lua_value_integer() {
1027 let lua = Lua::new();
1028 assert!(matches!(
1029 json_to_lua_value(&lua, serde_json::json!(42)).unwrap(),
1030 Value::Integer(42)
1031 ));
1032 }
1033
1034 #[test]
1035 fn test_json_to_lua_value_float() {
1036 let lua = Lua::new();
1037 assert!(matches!(
1038 json_to_lua_value(&lua, serde_json::json!(std::f64::consts::PI)).unwrap(),
1039 Value::Number(n) if (n - std::f64::consts::PI).abs() < f64::EPSILON
1040 ));
1041 }
1042
1043 #[test]
1044 fn test_json_to_lua_value_large_number() {
1045 let lua = Lua::new();
1046 let val = json_to_lua_value(&lua, serde_json::json!(1e200)).unwrap();
1048 assert!(matches!(val, Value::Number(_)));
1049 }
1050
1051 #[test]
1052 fn test_json_to_lua_value_string() {
1053 let lua = Lua::new();
1054 assert!(matches!(
1055 json_to_lua_value(&lua, serde_json::Value::String("hi".into())).unwrap(),
1056 Value::String(_)
1057 ));
1058 }
1059
1060 #[test]
1061 fn test_json_to_lua_value_array() {
1062 let lua = Lua::new();
1063 let val = json_to_lua_value(&lua, serde_json::json!([1, 2, 3])).unwrap();
1064 assert!(matches!(val, Value::Table(_)));
1065 }
1066
1067 #[test]
1068 fn test_json_to_lua_value_empty_array() {
1069 let lua = Lua::new();
1070 let val = json_to_lua_value(&lua, serde_json::json!([])).unwrap();
1071 assert!(matches!(val, Value::Table(_)));
1072 }
1073
1074 #[test]
1075 fn test_json_to_lua_value_object() {
1076 let lua = Lua::new();
1077 let val = json_to_lua_value(&lua, serde_json::json!({"a": 1})).unwrap();
1078 assert!(matches!(val, Value::Table(_)));
1079 }
1080
1081 #[test]
1082 fn test_json_to_lua_value_empty_object() {
1083 let lua = Lua::new();
1084 let val = json_to_lua_value(&lua, serde_json::json!({})).unwrap();
1085 assert!(matches!(val, Value::Table(_)));
1086 }
1087
1088 #[test]
1089 fn test_json_to_lua_value_nested() {
1090 let lua = Lua::new();
1091 let val = json_to_lua_value(&lua, serde_json::json!({"a": {"b": [1, 2, 3]}})).unwrap();
1092 assert!(matches!(val, Value::Table(_)));
1093 }
1094
1095 #[test]
1100 fn test_converge_result_construction() {
1101 let r = ConvergeResult {
1102 surviving_items: vec![serde_json::json!({"k": "v"})],
1103 findings: vec![default_finding()],
1104 rounds: 2,
1105 converged: true,
1106 round_stats: vec![RoundStats {
1107 round: 1,
1108 items_input: 3,
1109 findings_generated: 5,
1110 findings_survived: 4,
1111 findings_refuted: 1,
1112 approval_rate: 0.8,
1113 }],
1114 };
1115 assert_eq!(r.rounds, 2);
1116 assert!(r.converged);
1117 assert_eq!(r.surviving_items.len(), 1);
1118 assert_eq!(r.findings.len(), 1);
1119 assert_eq!(r.round_stats.len(), 1);
1120 }
1121
1122 #[test]
1123 fn test_converge_result_serialize_roundtrip() {
1124 let r = ConvergeResult {
1125 surviving_items: vec![],
1126 findings: vec![],
1127 rounds: 0,
1128 converged: true,
1129 round_stats: vec![],
1130 };
1131 let json = serde_json::to_string(&r).unwrap();
1132 let back: ConvergeResult = serde_json::from_str(&json).unwrap();
1133 assert_eq!(back.rounds, 0);
1134 assert!(back.converged);
1135 }
1136
1137 #[test]
1138 fn test_converge_result_debug_clone() {
1139 let r = ConvergeResult {
1140 surviving_items: vec![],
1141 findings: vec![],
1142 rounds: 1,
1143 converged: false,
1144 round_stats: vec![],
1145 };
1146 let _ = format!("{:?}", r);
1147 let _ = r.clone();
1148 }
1149
1150 #[test]
1151 fn test_round_stats_construction() {
1152 let s = RoundStats {
1153 round: 1,
1154 items_input: 10,
1155 findings_generated: 20,
1156 findings_survived: 15,
1157 findings_refuted: 5,
1158 approval_rate: 0.75,
1159 };
1160 assert_eq!(s.round, 1);
1161 assert_eq!(s.items_input, 10);
1162 assert_eq!(s.findings_generated, 20);
1163 assert_eq!(s.findings_survived, 15);
1164 assert_eq!(s.findings_refuted, 5);
1165 assert!((s.approval_rate - 0.75).abs() < f32::EPSILON);
1166 }
1167
1168 #[test]
1169 fn test_round_stats_serialize_roundtrip() {
1170 let s = RoundStats {
1171 round: 2,
1172 items_input: 5,
1173 findings_generated: 8,
1174 findings_survived: 3,
1175 findings_refuted: 5,
1176 approval_rate: 0.4,
1177 };
1178 let json = serde_json::to_string(&s).unwrap();
1179 let back: RoundStats = serde_json::from_str(&json).unwrap();
1180 assert_eq!(back.round, 2);
1181 assert!((back.approval_rate - 0.4).abs() < f32::EPSILON);
1182 }
1183
1184 #[tokio::test]
1189 async fn test_execute_convergence_empty_items() {
1190 let scheduler = converge_scheduler(Arc::new(NoOpBackend));
1191 let (ctx, _rx) = test_run_ctx(&scheduler);
1192
1193 let result = execute_convergence(
1194 vec![],
1195 "producer: {item}",
1196 "adversary: {finding}",
1197 ConvergeConfig::default(),
1198 &scheduler,
1199 &ctx,
1200 )
1201 .await
1202 .unwrap();
1203
1204 assert!(result.surviving_items.is_empty());
1205 assert!(result.findings.is_empty());
1206 assert_eq!(result.rounds, 0);
1207 assert!(result.converged);
1208 assert!(result.round_stats.is_empty());
1209 }
1210
1211 #[tokio::test]
1216 async fn test_execute_convergence_no_findings() {
1217 let mock = Arc::new(MockBackend::new(
1218 "mock",
1219 vec![MockBehavior::Success {
1220 output: serde_json::Value::Null,
1221 tokens: TokenUsage::default(),
1222 delay: Duration::from_millis(1),
1223 }],
1224 ));
1225 let scheduler = converge_scheduler(mock);
1226 let (ctx, _rx) = test_run_ctx(&scheduler);
1227
1228 let result = execute_convergence(
1229 vec![serde_json::json!({"key": "val"})],
1230 "producer: {item}",
1231 "adversary: {finding}",
1232 ConvergeConfig::default(),
1233 &scheduler,
1234 &ctx,
1235 )
1236 .await
1237 .unwrap();
1238
1239 assert_eq!(result.rounds, 0);
1242 assert!(!result.converged);
1243 assert!(result.findings.is_empty());
1244 assert_eq!(
1245 result.surviving_items,
1246 vec![serde_json::json!({"key": "val"})]
1247 );
1248 assert!(result.round_stats.is_empty());
1249 }
1250
1251 #[tokio::test]
1256 async fn test_execute_convergence_findings_survive() {
1257 let findings = vec![default_finding()];
1258 let backend = Arc::new(FindingsAlwaysBackend { findings });
1259 let scheduler = converge_scheduler(backend);
1260 let (ctx, _rx) = test_run_ctx(&scheduler);
1261
1262 let result = execute_convergence(
1263 vec![serde_json::json!({"input": "item"})],
1264 "producer: {item}",
1265 "adversary: {finding}",
1266 ConvergeConfig::default(),
1267 &scheduler,
1268 &ctx,
1269 )
1270 .await
1271 .unwrap();
1272
1273 assert_eq!(result.rounds, 2);
1276 assert!(result.converged);
1277 assert!(!result.findings.is_empty());
1278 assert_eq!(result.round_stats.len(), 2);
1279 for rs in &result.round_stats {
1281 assert_eq!(rs.findings_generated, 1);
1282 assert_eq!(rs.findings_survived, 1);
1283 assert_eq!(rs.findings_refuted, 0);
1284 }
1285 }
1286
1287 #[tokio::test]
1292 async fn test_execute_convergence_non_adversarial() {
1293 let findings = vec![default_finding()];
1294 let backend = Arc::new(FindingsAlwaysBackend { findings });
1295 let scheduler = converge_scheduler(backend);
1296 let (ctx, _rx) = test_run_ctx(&scheduler);
1297
1298 let config = ConvergeConfig {
1299 adversarial: false,
1300 ..ConvergeConfig::default()
1301 };
1302
1303 let result = execute_convergence(
1304 vec![serde_json::json!({"x": 1})],
1305 "producer: {item}",
1306 "adversary: {finding}",
1307 config,
1308 &scheduler,
1309 &ctx,
1310 )
1311 .await
1312 .unwrap();
1313
1314 assert_eq!(result.rounds, 2);
1316 assert!(result.converged);
1317 assert!(!result.findings.is_empty());
1318 }
1319
1320 #[tokio::test]
1325 async fn test_execute_convergence_all_refuted() {
1326 let findings = vec![default_finding()];
1327 let backend = Arc::new(FindingsRefutingBackend { findings });
1328 let scheduler = converge_scheduler(backend);
1329 let (ctx, _rx) = test_run_ctx(&scheduler);
1330
1331 let result = execute_convergence(
1332 vec![serde_json::json!({"x": 1})],
1333 "producer: {item}",
1334 "adversary: {finding}",
1335 ConvergeConfig::default(),
1336 &scheduler,
1337 &ctx,
1338 )
1339 .await
1340 .unwrap();
1341
1342 assert_eq!(result.rounds, 1);
1344 assert!(result.converged);
1345 assert_eq!(result.round_stats.len(), 1);
1346 assert_eq!(result.round_stats[0].findings_generated, 1);
1347 assert_eq!(result.round_stats[0].findings_survived, 0);
1348 assert_eq!(result.round_stats[0].findings_refuted, 1);
1349 }
1350
1351 #[tokio::test]
1356 async fn test_execute_convergence_no_adversarial_findings_generated() {
1357 let findings = vec![default_finding()];
1358 let backend = Arc::new(FindingsAlwaysBackend { findings });
1359 let scheduler = converge_scheduler(backend);
1360 let (ctx, _rx) = test_run_ctx(&scheduler);
1361
1362 let config = ConvergeConfig {
1363 adversarial: false,
1364 max_rounds: 1,
1365 ..ConvergeConfig::default()
1366 };
1367
1368 let result = execute_convergence(
1369 vec![serde_json::json!({"x": 1})],
1370 "producer: {item}",
1371 "adversary: {finding}",
1372 config,
1373 &scheduler,
1374 &ctx,
1375 )
1376 .await
1377 .unwrap();
1378
1379 assert_eq!(result.rounds, 1);
1380 assert!(!result.converged);
1382 assert_eq!(result.findings.len(), 1);
1383 }
1384
1385 #[test]
1390 fn test_register_converge_sdk_empty_items() {
1391 let rt = tokio::runtime::Runtime::new().unwrap();
1392 let lua = Lua::new();
1393 let scheduler = converge_scheduler(Arc::new(NoOpBackend));
1394 let run_id = Uuid::now_v7();
1395 let (tx, _rx2) = broadcast::channel(64);
1396 let run_ctx = RunContext {
1397 run_id,
1398 cancel: CancellationToken::new(),
1399 events: tx,
1400 };
1401 let report_sink: ReportSink = Arc::new(std::sync::Mutex::new(None));
1402 let handle = rt.handle().clone();
1403 let cx = SdkContext::new(run_ctx, scheduler, report_sink, None, handle);
1404
1405 register_converge_sdk(&lua, &cx).unwrap();
1406
1407 let globals = lua.globals();
1408 let converge: mlua::Function = globals.get("converge").unwrap();
1409
1410 let items = lua.create_table().unwrap();
1411 let options = lua.create_table().unwrap();
1412 let result: mlua::Table = converge.call((items, options)).unwrap();
1413
1414 let converged: bool = result.get("converged").unwrap();
1415 assert!(converged);
1416 let rounds: u32 = result.get("rounds").unwrap();
1417 assert_eq!(rounds, 0);
1418 let surviving: mlua::Table = result.get("surviving").unwrap();
1419 assert_eq!(surviving.len().unwrap(), 0);
1420 let findings: mlua::Table = result.get("findings").unwrap();
1421 assert_eq!(findings.len().unwrap(), 0);
1422 }
1423
1424 #[test]
1425 fn test_register_converge_sdk_with_items() {
1426 let rt = tokio::runtime::Runtime::new().unwrap();
1427 let lua = Lua::new();
1428 let backend = Arc::new(FindingsAlwaysBackend {
1429 findings: vec![default_finding()],
1430 });
1431 let scheduler = converge_scheduler(backend);
1432 let run_id = Uuid::now_v7();
1433 let _rx = scheduler.init_run(run_id, 64);
1434 let (tx, _rx2) = broadcast::channel(64);
1435 let run_ctx = RunContext {
1436 run_id,
1437 cancel: CancellationToken::new(),
1438 events: tx,
1439 };
1440 let report_sink: ReportSink = Arc::new(std::sync::Mutex::new(None));
1441 let handle = rt.handle().clone();
1442 let cx = SdkContext::new(run_ctx, scheduler, report_sink, None, handle);
1443
1444 register_converge_sdk(&lua, &cx).unwrap();
1445
1446 let globals = lua.globals();
1447 let converge: mlua::Function = globals.get("converge").unwrap();
1448
1449 let items = lua.create_table().unwrap();
1450 items.set(1, "test item").unwrap();
1451 let options = lua.create_table().unwrap();
1452 let result: mlua::Table = converge.call((items, options)).unwrap();
1453
1454 let converged: bool = result.get("converged").unwrap();
1455 assert!(converged);
1456 let rounds: u32 = result.get("rounds").unwrap();
1457 assert!(rounds > 0);
1458 let findings: mlua::Table = result.get("findings").unwrap();
1459 assert!(findings.len().unwrap() > 0);
1460 }
1461
1462 #[test]
1467 fn test_producer_stats() {
1468 let s = ProducerStats {
1469 items_processed: 5,
1470 agents_run: 10,
1471 findings_generated: 20,
1472 };
1473 assert_eq!(s.items_processed, 5);
1474 assert_eq!(s.agents_run, 10);
1475 assert_eq!(s.findings_generated, 20);
1476 }
1477
1478 #[test]
1479 fn test_producer_stats_default() {
1480 let s = ProducerStats::default();
1481 assert_eq!(s.items_processed, 0);
1482 assert_eq!(s.agents_run, 0);
1483 assert_eq!(s.findings_generated, 0);
1484 }
1485
1486 #[test]
1487 fn test_vote_stats() {
1488 let s = VoteStats { approval_rate: 0.5 };
1489 assert!((s.approval_rate - 0.5).abs() < f32::EPSILON);
1490 }
1491
1492 #[test]
1493 fn test_vote_stats_default() {
1494 let s = VoteStats::default();
1495 assert!((s.approval_rate - 0.0).abs() < f32::EPSILON);
1496 }
1497
1498 struct NoOpBackend;
1505
1506 #[async_trait]
1507 impl AgentBackend for NoOpBackend {
1508 fn id(&self) -> &'static str {
1509 "noop"
1510 }
1511 fn capabilities(&self) -> AgentCapabilities {
1512 AgentCapabilities::default()
1513 }
1514 fn as_any(&self) -> &dyn std::any::Any {
1515 self
1516 }
1517 async fn run(
1518 &self,
1519 task: AgentTask,
1520 _ctx: RunContext,
1521 ) -> Result<AgentResult, BackendError> {
1522 Ok(AgentResult {
1523 agent_id: task.agent_id,
1524 status: AgentStatus::Ok,
1525 output: serde_json::Value::Null,
1526 findings: vec![],
1527 tokens_used: TokenUsage::default(),
1528 artifacts: vec![],
1529 logs: LogRef::default(),
1530 })
1531 }
1532 }
1533
1534 struct FindingsAlwaysBackend {
1537 findings: Vec<Finding>,
1538 }
1539
1540 #[async_trait]
1541 impl AgentBackend for FindingsAlwaysBackend {
1542 fn id(&self) -> &'static str {
1543 "findings-always"
1544 }
1545 fn capabilities(&self) -> AgentCapabilities {
1546 AgentCapabilities::default()
1547 }
1548 fn as_any(&self) -> &dyn std::any::Any {
1549 self
1550 }
1551 async fn run(
1552 &self,
1553 task: AgentTask,
1554 _ctx: RunContext,
1555 ) -> Result<AgentResult, BackendError> {
1556 Ok(AgentResult {
1557 agent_id: task.agent_id,
1558 status: AgentStatus::Ok,
1559 output: serde_json::Value::Null,
1560 findings: self.findings.clone(),
1561 tokens_used: TokenUsage::default(),
1562 artifacts: vec![],
1563 logs: LogRef::default(),
1564 })
1565 }
1566 }
1567
1568 struct FindingsRefutingBackend {
1571 findings: Vec<Finding>,
1572 }
1573
1574 #[async_trait]
1575 impl AgentBackend for FindingsRefutingBackend {
1576 fn id(&self) -> &'static str {
1577 "findings-refute"
1578 }
1579 fn capabilities(&self) -> AgentCapabilities {
1580 AgentCapabilities::default()
1581 }
1582 fn as_any(&self) -> &dyn std::any::Any {
1583 self
1584 }
1585 async fn run(
1586 &self,
1587 task: AgentTask,
1588 _ctx: RunContext,
1589 ) -> Result<AgentResult, BackendError> {
1590 Ok(AgentResult {
1591 agent_id: task.agent_id,
1592 status: AgentStatus::Error,
1593 output: serde_json::Value::Null,
1594 findings: self.findings.clone(),
1595 tokens_used: TokenUsage::default(),
1596 artifacts: vec![],
1597 logs: LogRef::default(),
1598 })
1599 }
1600 }
1601}